Skip to content

Engines

Every engine subclasses DelineationEngine and is resolved by name with get_engine. See Delineation engines for an overview.

engines

Delineation engines for agricultural field boundary detection.

Each engine wraps a different model or approach for extracting field boundary polygons from satellite imagery or embeddings. Engine modules are imported lazily by :func:get_engine, so importing this package does not pull in torch or other optional dependencies.

DelineationEngine

Bases: ABC

Abstract base class for delineation engines.

Subclasses must implement :meth:delineate. They may attach JSON-serialisable run metadata (backend, model id, weights repository, revision, sha256, thresholds, window dates, ...) to the returned frame as gdf.attrs["engine_meta"]; the pipeline copies it into the provenance record.

Source code in agribound/engines/base.py
class DelineationEngine(ABC):
    """Abstract base class for delineation engines.

    Subclasses must implement :meth:`delineate`. They may attach
    JSON-serialisable run metadata (backend, model id, weights repository,
    revision, sha256, thresholds, window dates, ...) to the returned frame as
    ``gdf.attrs["engine_meta"]``; the pipeline copies it into the provenance
    record.
    """

    name: str = "base"
    supported_sources: list[str] = []
    requires_bands: list[str] = []

    @abstractmethod
    def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
        """Run field boundary delineation on a raster file.

        Parameters
        ----------
        raster_path : str
            Path to the input GeoTIFF (composite or local file).
        config : AgriboundConfig
            Pipeline configuration.

        Returns
        -------
        geopandas.GeoDataFrame
            Field boundary polygons with at minimum a ``geometry`` column.
            ``gdf.attrs["engine_meta"]`` may hold engine metadata.
        """

    def validate_input(self, raster_path: str, config: AgriboundConfig) -> None:
        """Validate that the input raster is compatible with this engine.

        Checks that the raster has enough bands for the engine's
        ``requires_bands`` (class attribute, falling back to the registry
        entry): at least the highest 1-based index those canonical bands map
        to for the configured source, with ``config.bands`` taking precedence
        (:func:`get_canonical_band_indices`). For local rasters without
        ``config.bands`` the indices are positional, so this is
        ``len(requires_bands)``.

        Parameters
        ----------
        raster_path : str
            Path to the input raster.
        config : AgriboundConfig
            Pipeline configuration.

        Raises
        ------
        ValueError
            If the input is incompatible.
        """
        from agribound.io.raster import get_raster_info

        required = list(self.requires_bands) or list(
            ENGINE_REGISTRY.get(self.name, {}).get("requires_bands", [])
        )
        info = get_raster_info(raster_path)
        n_needed = 0
        if required:
            try:
                n_needed = max(
                    get_canonical_band_indices(config.source, required, bands=config.bands)
                )
            except ValueError as exc:
                raise ValueError(
                    f"Engine {self.name!r} requires bands {required}, which source "
                    f"{config.source!r} does not provide: {exc}"
                ) from exc
        if n_needed > 0 and info.count < n_needed:
            override = f" with bands={config.bands}" if config.bands else ""
            raise ValueError(
                f"Engine {self.name!r} requires bands {required}{override}, i.e. at least "
                f"{n_needed} bands, but the raster has {info.count} bands."
            )

    @classmethod
    def prefetch(cls, config: AgriboundConfig) -> list[str]:
        """Download model weights so that inference can run offline.

        Engines that load remote weights override this and return the local
        paths (or cache directories) they populated. The base implementation
        downloads nothing.

        Parameters
        ----------
        config : AgriboundConfig
            Pipeline configuration (engine parameters select the model).

        Returns
        -------
        list[str]
            Local paths of the downloaded artefacts (empty here).
        """
        return []

delineate abstractmethod

delineate(raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame

Run field boundary delineation on a raster file.

Parameters:

Name Type Description Default
raster_path str

Path to the input GeoTIFF (composite or local file).

required
config AgriboundConfig

Pipeline configuration.

required

Returns:

Type Description
GeoDataFrame

Field boundary polygons with at minimum a geometry column. gdf.attrs["engine_meta"] may hold engine metadata.

Source code in agribound/engines/base.py
@abstractmethod
def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
    """Run field boundary delineation on a raster file.

    Parameters
    ----------
    raster_path : str
        Path to the input GeoTIFF (composite or local file).
    config : AgriboundConfig
        Pipeline configuration.

    Returns
    -------
    geopandas.GeoDataFrame
        Field boundary polygons with at minimum a ``geometry`` column.
        ``gdf.attrs["engine_meta"]`` may hold engine metadata.
    """

validate_input

validate_input(raster_path: str, config: AgriboundConfig) -> None

Validate that the input raster is compatible with this engine.

Checks that the raster has enough bands for the engine's requires_bands (class attribute, falling back to the registry entry): at least the highest 1-based index those canonical bands map to for the configured source, with config.bands taking precedence (:func:get_canonical_band_indices). For local rasters without config.bands the indices are positional, so this is len(requires_bands).

Parameters:

Name Type Description Default
raster_path str

Path to the input raster.

required
config AgriboundConfig

Pipeline configuration.

required

Raises:

Type Description
ValueError

If the input is incompatible.

Source code in agribound/engines/base.py
def validate_input(self, raster_path: str, config: AgriboundConfig) -> None:
    """Validate that the input raster is compatible with this engine.

    Checks that the raster has enough bands for the engine's
    ``requires_bands`` (class attribute, falling back to the registry
    entry): at least the highest 1-based index those canonical bands map
    to for the configured source, with ``config.bands`` taking precedence
    (:func:`get_canonical_band_indices`). For local rasters without
    ``config.bands`` the indices are positional, so this is
    ``len(requires_bands)``.

    Parameters
    ----------
    raster_path : str
        Path to the input raster.
    config : AgriboundConfig
        Pipeline configuration.

    Raises
    ------
    ValueError
        If the input is incompatible.
    """
    from agribound.io.raster import get_raster_info

    required = list(self.requires_bands) or list(
        ENGINE_REGISTRY.get(self.name, {}).get("requires_bands", [])
    )
    info = get_raster_info(raster_path)
    n_needed = 0
    if required:
        try:
            n_needed = max(
                get_canonical_band_indices(config.source, required, bands=config.bands)
            )
        except ValueError as exc:
            raise ValueError(
                f"Engine {self.name!r} requires bands {required}, which source "
                f"{config.source!r} does not provide: {exc}"
            ) from exc
    if n_needed > 0 and info.count < n_needed:
        override = f" with bands={config.bands}" if config.bands else ""
        raise ValueError(
            f"Engine {self.name!r} requires bands {required}{override}, i.e. at least "
            f"{n_needed} bands, but the raster has {info.count} bands."
        )

prefetch classmethod

prefetch(config: AgriboundConfig) -> list[str]

Download model weights so that inference can run offline.

Engines that load remote weights override this and return the local paths (or cache directories) they populated. The base implementation downloads nothing.

Parameters:

Name Type Description Default
config AgriboundConfig

Pipeline configuration (engine parameters select the model).

required

Returns:

Type Description
list[str]

Local paths of the downloaded artefacts (empty here).

Source code in agribound/engines/base.py
@classmethod
def prefetch(cls, config: AgriboundConfig) -> list[str]:
    """Download model weights so that inference can run offline.

    Engines that load remote weights override this and return the local
    paths (or cache directories) they populated. The base implementation
    downloads nothing.

    Parameters
    ----------
    config : AgriboundConfig
        Pipeline configuration (engine parameters select the model).

    Returns
    -------
    list[str]
        Local paths of the downloaded artefacts (empty here).
    """
    return []

get_canonical_band_indices

get_canonical_band_indices(source: str, canonical_names: list[str], bands: dict[str, int] | None = None) -> list[int]

Get 1-based raster band indices for canonical band names.

Looks up each canonical name ("R", "G", "B", "NIR", "NIR_NARROW", "SWIR1", "SWIR2") in the source registry and returns the corresponding 1-based band index in the composite written by the source's builder.

Parameters:

Name Type Description Default
source str

Satellite source name.

required
canonical_names list[str]

Canonical band names to look up (e.g. ["R", "G", "B"]).

required
bands dict[str, int] or None

Optional explicit mapping of canonical names to 1-based indices (for example AgriboundConfig.bands). When given it takes precedence over the registry for every name it contains.

None

Returns:

Type Description
list[int]

1-based band indices in the composite raster.

Raises:

Type Description
ValueError

If the source is unknown or a canonical band is not available.

Notes

For source="local" without an explicit mapping the indices are positional (1, 2, 3, ... in the order requested), i.e. the local file is assumed to store the requested bands first and in that order.

Source code in agribound/engines/base.py
def get_canonical_band_indices(
    source: str,
    canonical_names: list[str],
    bands: dict[str, int] | None = None,
) -> list[int]:
    """Get 1-based raster band indices for canonical band names.

    Looks up each canonical name (``"R"``, ``"G"``, ``"B"``, ``"NIR"``,
    ``"NIR_NARROW"``, ``"SWIR1"``, ``"SWIR2"``) in the source registry and
    returns the corresponding 1-based band index in the composite written by
    the source's builder.

    Parameters
    ----------
    source : str
        Satellite source name.
    canonical_names : list[str]
        Canonical band names to look up (e.g. ``["R", "G", "B"]``).
    bands : dict[str, int] or None
        Optional explicit mapping of canonical names to 1-based indices (for
        example ``AgriboundConfig.bands``). When given it takes precedence
        over the registry for every name it contains.

    Returns
    -------
    list[int]
        1-based band indices in the composite raster.

    Raises
    ------
    ValueError
        If the source is unknown or a canonical band is not available.

    Notes
    -----
    For ``source="local"`` without an explicit mapping the indices are
    positional (``1, 2, 3, ...`` in the order requested), i.e. the local file
    is assumed to store the requested bands first and in that order.
    """
    info = SOURCE_REGISTRY.get(source)
    if info is None:
        raise ValueError(f"Unknown source {source!r}")

    all_bands = info.get("all_bands")
    canonical = info.get("canonical_bands") or {}
    bands = dict(bands or {})

    if all_bands is None and not bands:
        # Local source -- positional (1, 2, 3, ...)
        return list(range(1, len(canonical_names) + 1))

    indices = []
    for position, name in enumerate(canonical_names):
        if name in bands:
            idx = int(bands[name])
            if idx < 1:
                raise ValueError(f"Band index for {name!r} must be >= 1, got {idx}")
            indices.append(idx)
            continue
        if all_bands is None:
            # Local source with a partial mapping: remaining names positional.
            indices.append(position + 1)
            continue
        native = canonical.get(name)
        if native is None:
            known = [n for n in CANONICAL_BAND_NAMES if n in canonical]
            raise ValueError(
                f"Canonical band {name!r} not defined for source {source!r}. Available: {known}"
            )
        indices.append(all_bands.index(native) + 1)  # 1-based
    return indices

get_engine

get_engine(engine_name: str) -> DelineationEngine

Factory function to get a delineation engine instance by name.

Parameters:

Name Type Description Default
engine_name str

Engine name (e.g. "delineate-anything", "ftw").

required

Returns:

Type Description
DelineationEngine

Engine instance.

Raises:

Type Description
ValueError

If the engine name is not recognised.

Source code in agribound/engines/base.py
def get_engine(engine_name: str) -> DelineationEngine:
    """Factory function to get a delineation engine instance by name.

    Parameters
    ----------
    engine_name : str
        Engine name (e.g. ``"delineate-anything"``, ``"ftw"``).

    Returns
    -------
    DelineationEngine
        Engine instance.

    Raises
    ------
    ValueError
        If the engine name is not recognised.
    """
    return get_engine_class(engine_name)()

get_engine_class

get_engine_class(engine_name: str) -> type[DelineationEngine]

Import and return the engine class for engine_name without instantiating it.

Parameters:

Name Type Description Default
engine_name str

Engine name (e.g. "delineate-anything").

required

Returns:

Type Description
type[DelineationEngine]

Engine class.

Raises:

Type Description
ValueError

If the engine name is not recognised.

Source code in agribound/engines/base.py
def get_engine_class(engine_name: str) -> type[DelineationEngine]:
    """Import and return the engine class for *engine_name* without instantiating it.

    Parameters
    ----------
    engine_name : str
        Engine name (e.g. ``"delineate-anything"``).

    Returns
    -------
    type[DelineationEngine]
        Engine class.

    Raises
    ------
    ValueError
        If the engine name is not recognised.
    """
    key = str(engine_name).lower().strip()
    target = ENGINE_CLASSES.get(key)
    if target is None or key not in ENGINE_REGISTRY:
        raise ValueError(f"Unknown engine {engine_name!r}. Available: {list(ENGINE_REGISTRY)}")
    module_name, _, class_name = target.partition(":")
    module = importlib.import_module(module_name)
    return getattr(module, class_name)

list_engines

list_engines() -> dict[str, dict[str, Any]]

List all delineation engines and their metadata.

Returns:

Type Description
dict[str, dict]

Deep copy of :data:ENGINE_REGISTRY (safe to mutate).

Examples:

>>> from agribound import list_engines
>>> for name, info in list_engines().items():
...     print(name, info["approach"])
Source code in agribound/registry.py
def list_engines() -> dict[str, dict[str, Any]]:
    """List all delineation engines and their metadata.

    Returns
    -------
    dict[str, dict]
        Deep copy of :data:`ENGINE_REGISTRY` (safe to mutate).

    Examples
    --------
    >>> from agribound import list_engines
    >>> for name, info in list_engines().items():
    ...     print(name, info["approach"])
    """
    return copy.deepcopy(ENGINE_REGISTRY)

Delineate-Anything

delineate_anything

Delineate-Anything engine (YOLO11-seg instance segmentation of field boundaries).

Models

Weights come from the Hugging Face repository MykolaL/DelineateAnything at pinned revisions; the SHA-256 of every downloaded file is checked against :data:DA_MODELS before it is used.

  • large_v2 (default): DelineateAnythingv2.pt, Delineate Anything v2, YOLO11x-seg trained on FBIS-73M; default confidence 0.15; FTW registry name DelineateAnythingV2.
  • large: DelineateAnything.pt, YOLO11x-seg trained on FBIS-22M; default confidence 0.005; FTW name DelineateAnything.
  • small: DelineateAnything-S.pt, YOLO11n-seg trained on FBIS-22M; default confidence 0.005; FTW name DelineateAnything-S.

The default confidences are those of the upstream conf_sample.yaml: 0.15 for large_v2 (Lavreniuk/Delineate-Anything a6f30b2) and 0.005 for the v1 models (the v1-era sample configuration). Select a model with engine_params["da_model"] (a key or one of the aliases "DelineateAnythingV2", "DelineateAnything", "DelineateAnything-S"); the legacy engine_params["model_size"] ("large"/"small") selects the v1 models.

Backends

engine_params["backend"] chooses the implementation explicitly (default "native"). There is no automatic fallback between backends: a backend that cannot run raises an error that says what is missing.

"native" Agribound's own tiled Ultralytics inference. It reproduces the preprocessing of the reference Delineate-Anything pipeline (DelAnyFlow, upstream methods/main): a scene-level per-band 1-99 percentile stretch to uint8, computed on valid, strictly positive pixels sampled (nearest neighbour, full-resolution data, never overviews) on a grid of at most 4096 px per side (uint8 rasters are used unchanged); tiles of 512 native pixels when the ground sampling distance (GSD) is below 4 m, else 256 native pixels upsampled 2x with bicubic interpolation, so the model input is always 512 x 512 (super_resolution = 1, 2 or 4 overrides the factor); 50 % tile overlap starting half a tile before the raster origin, as the upstream ExecutionPlanner does; BGR channel order for NumPy input to Ultralytics; retina_masks=True; masks cast to float before the upstream 3 x 3 erode / dilate / dilate / erode morphology (Ultralytics >= 8.3.217 returns uint8 masks, on which the upstream negation trick would turn the erosion into a dilation); FP16 on GPU/MPS. Each detection becomes the largest polygon of its mask, clipped to valid pixels, and is flagged when it touches an interior tile edge (a tile-cut detection). The detections of all tiles are then combined at polygon level, with duplicates defined as IoU >= dedup_iou or intersection >= dedup_containment of the smaller polygon: (1) tile-cut duplicates of one another are merged into their union (:func:merge_tile_pieces; merge_tile_pieces=False skips it), so a field too large to be complete in any tile is rebuilt from its pieces (the polygon-level counterpart of DelAnyFlow's merging of fields that touch a tile border); pieces that duplicate a complete (not tile-cut), at least as large detection are not merged; (2) greedy non-maximum suppression visits the polygons by higher confidence, then larger area, and drops every polygon that duplicates one already kept, except that a complete detection is visited before each tile-cut duplicate that is not larger than it (:func:deduplicate_detections); (3) remaining overlaps go to the polygon visited first in that order (:func:resolve_overlaps; resolve_overlaps=False keeps them). This is simpler than DelAnyFlow's raster-level region merging (which, for example, lets smaller fields carve their area out of larger ones), so results are close to, but not identical with, the "reference" backend. NMS uses IoU 0.3 by default, the value of the authors' openEO UDP (openeo_udp/udf/delineate_onnx.py); the upstream execute() keeps Ultralytics' default (0.7). The raster is read tile by tile. "reference" Runs the upstream DelAnyFlow pipeline (methods.main.inference.execute) in a subprocess, from a Delineate-Anything checkout given by engine_params["da_repo"] or the AGRIBOUND_DA_REPO environment variable. Requires the GDAL Python bindings (osgeo) and a checkout that contains the uint8-mask fix (upstream commit 34eddf7 or later). "ftw" ftw_tools.inference.inference.run_instance_segmentation. FTW's wrapper divides the first three bands by 3000 (Sentinel-2 L2A units), clips to [0, 1] and resizes bilinearly, so this backend accepts only reflectance_x10000 composites. large_v2 needs an ftw-tools build whose MODEL_REGISTRY contains DelineateAnythingV2 (ftw-baselines main at fa86d4a or later; not in ftw-tools 2.0.0b5).

Engine parameters

All optional. A Delineate-Anything parameter that the selected backend cannot honour raises :class:ValueError, as do the names confidence and minimal_confidence (the confidence is conf_threshold for every backend); parameters not listed here are left to other pipeline stages.

  • All backends: backend; da_model/model_size (large_v2); conf_threshold (per model; the ftw backend uses 0.15 for v2 and ftw-tools' 0.05 for v1); batch_size (tiles per forward pass, 4).
  • native and reference: checkpoint_path (fine-tuned YOLO weights; set by the pipeline after fine-tuning); super_resolution (1, 2 or 4; default automatic); tile_step (fraction of the tile, 0.5); half (FP16 on GPU/MPS, True).
  • native and ftw: iou_threshold (NMS IoU, 0.3); max_detections (per tile, 300).
  • native: dedup_iou (0.3), dedup_containment (0.8), merge_tile_pieces (True) and resolve_overlaps (True).
  • reference: da_repo; min_hole_area_m2 (holes smaller than this are filled, 2500 m², the upstream conf_sample.yaml value).
  • ftw: patch_size (256; a multiple of 32 smaller than the raster's smaller side), resize_factor (2), padding (FTW default), close_interiors (True), simplify (FTW simplification tolerance, applied in EPSG:6933 coordinates, i.e. metres that are exact only near 30° latitude; 0), max_size (m², None), overlap_iou_threshold (0.3), overlap_contain_threshold (0.8), value_scale (for local rasters).

config.min_field_area_m2 is applied as an absolute area in m², computed in the equal-area EPSG:6933, by every backend (the reference pipeline's automatic_area_scale is disabled; upstream and ftw-tools compute areas in EPSG:6933 too). Holes: the native backend keeps holes of any size, the reference backend fills holes smaller than min_hole_area_m2 and the ftw backend fills all holes while close_interiors is True. The returned frame carries gdf.attrs["engine_meta"] (backend, model key, weights repository, revision and SHA-256, thresholds, super-resolution factor, tile size, device, pixel size, ...); the native backend also returns a confidence column. The published models were trained on 0.25-10 m imagery: for rasters outside that range (e.g. 30 m Landsat/HLS) a WARNING is logged and engine_meta["gsd_outside_training_range"] is True (:func:gsd_outside_training_range).

References

Lavreniuk, M., et al. (2025). Delineate Anything: Resolution-Agnostic Field Boundary Delineation on Satellite Imagery. European Conference on Artificial Intelligence (ECAI 2025). arXiv:2504.02534.

Lavreniuk, M., et al. (2025). Delineate Anything Flow: Fast, Country-Level Field Boundary Detection from Any Source. arXiv:2511.13417.

Lavreniuk, M., et al. (2026). Delineate Anything v2: A Global Foundation Model for Field Delineation. European Conference on Computer Vision Workshops (ECCVW 2026), GAIA workshop. arXiv:2607.19069.

Model code and weights are AGPL-3.0; Ultralytics is AGPL-3.0.

DA_MODELS module-attribute

DA_MODELS: dict[str, DAModel] = {'large_v2': DAModel(key='large_v2', filename='DelineateAnythingv2.pt', revision='369d0b4c44cf9bec2bd3a27bc81810cadd2c963e', sha256='46700b8a279b07922953a11adaeb5e658d9a2384b6334c8e0a3090886218915a', size_bytes=124747297, default_conf=0.15, ftw_name='DelineateAnythingV2', architecture='YOLO11x-seg', training_data='FBIS-73M'), 'large': DAModel(key='large', filename='DelineateAnything.pt', revision='029e9a94c6abc51c67cebdc9b9a9b6c1ac2b1187', sha256='e3dcda35780083aeaefe9425b73b15a30561cdc12c43277041d04f9e88ede029', size_bytes=124746842, default_conf=0.005, ftw_name='DelineateAnything', architecture='YOLO11x-seg', training_data='FBIS-22M'), 'small': DAModel(key='small', filename='DelineateAnything-S.pt', revision='029e9a94c6abc51c67cebdc9b9a9b6c1ac2b1187', sha256='5463cdfb73690fc506035e4f7dce26a4c06af6ef4d207570513110c4b879d643', size_bytes=17635629, default_conf=0.005, ftw_name='DelineateAnything-S', architecture='YOLO11n-seg', training_data='FBIS-22M')}

Pinned Delineate-Anything checkpoints (see the module docstring).

DelineateAnythingEngine

Bases: DelineationEngine

Field boundary delineation with Delineate-Anything (see the module docstring).

Source code in agribound/engines/delineate_anything.py
class DelineateAnythingEngine(DelineationEngine):
    """Field boundary delineation with Delineate-Anything (see the module docstring)."""

    name = "delineate-anything"
    supported_sources = list(_REGISTRY_ENTRY["supported_sources"])
    requires_bands = list(_REGISTRY_ENTRY["requires_bands"])

    def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
        """Run Delineate-Anything on a raster.

        Parameters
        ----------
        raster_path : str
            Input GeoTIFF (composite or local file).
        config : AgriboundConfig
            Pipeline configuration; ``engine_params`` select the backend and
            model (module docstring).

        Returns
        -------
        geopandas.GeoDataFrame
            Field polygons in the raster CRS with
            ``gdf.attrs["engine_meta"]``.
        """
        self.validate_input(raster_path, config)
        opts = DAOptions.from_engine_params(config.engine_params)
        rgb = get_canonical_band_indices(config.source, ["R", "G", "B"], bands=config.bands)
        if opts.backend == "native":
            return _run_native(raster_path, config, opts, rgb)
        if opts.backend == "reference":
            return _run_reference(raster_path, config, opts, rgb)
        return _run_ftw(raster_path, config, opts, rgb)

    @classmethod
    def prefetch(cls, config: AgriboundConfig) -> list[str]:
        """Download the weights selected by ``config.engine_params``.

        ``native``/``reference``: the pinned Hugging Face file (SHA-256
        checked), unless ``checkpoint_path`` is set (then that file is
        returned if it exists). ``ftw``: ftw-tools' checkpoint URL, which
        Ultralytics resolves relative to the current working directory, is
        downloaded into the current working directory, so the later run must
        start from the same directory to find it offline.

        Returns
        -------
        list[str]
            Local paths of the weights.
        """
        opts = DAOptions.from_engine_params(config.engine_params)
        if opts.checkpoint_path:
            path = Path(opts.checkpoint_path).expanduser()
            if not path.is_file():
                raise FileNotFoundError(f"checkpoint_path does not exist: {path}")
            return [str(path.resolve())]
        if opts.backend in ("native", "reference"):
            return [download_da_weights(opts.model_key)]
        url = _ftw_checkpoint_url(opts.model.ftw_name)
        if url is None:
            raise RuntimeError(
                f"The installed ftw-tools has no checkpoint URL for {opts.model.ftw_name!r}."
            )
        from urllib.parse import unquote, urlparse

        import torch

        target = Path.cwd() / Path(unquote(urlparse(url).path)).name
        if not target.is_file():
            torch.hub.download_url_to_file(url, str(target), progress=True)
        return [str(target)]
delineate
delineate(raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame

Run Delineate-Anything on a raster.

Parameters:

Name Type Description Default
raster_path str

Input GeoTIFF (composite or local file).

required
config AgriboundConfig

Pipeline configuration; engine_params select the backend and model (module docstring).

required

Returns:

Type Description
GeoDataFrame

Field polygons in the raster CRS with gdf.attrs["engine_meta"].

Source code in agribound/engines/delineate_anything.py
def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
    """Run Delineate-Anything on a raster.

    Parameters
    ----------
    raster_path : str
        Input GeoTIFF (composite or local file).
    config : AgriboundConfig
        Pipeline configuration; ``engine_params`` select the backend and
        model (module docstring).

    Returns
    -------
    geopandas.GeoDataFrame
        Field polygons in the raster CRS with
        ``gdf.attrs["engine_meta"]``.
    """
    self.validate_input(raster_path, config)
    opts = DAOptions.from_engine_params(config.engine_params)
    rgb = get_canonical_band_indices(config.source, ["R", "G", "B"], bands=config.bands)
    if opts.backend == "native":
        return _run_native(raster_path, config, opts, rgb)
    if opts.backend == "reference":
        return _run_reference(raster_path, config, opts, rgb)
    return _run_ftw(raster_path, config, opts, rgb)
prefetch classmethod
prefetch(config: AgriboundConfig) -> list[str]

Download the weights selected by config.engine_params.

native/reference: the pinned Hugging Face file (SHA-256 checked), unless checkpoint_path is set (then that file is returned if it exists). ftw: ftw-tools' checkpoint URL, which Ultralytics resolves relative to the current working directory, is downloaded into the current working directory, so the later run must start from the same directory to find it offline.

Returns:

Type Description
list[str]

Local paths of the weights.

Source code in agribound/engines/delineate_anything.py
@classmethod
def prefetch(cls, config: AgriboundConfig) -> list[str]:
    """Download the weights selected by ``config.engine_params``.

    ``native``/``reference``: the pinned Hugging Face file (SHA-256
    checked), unless ``checkpoint_path`` is set (then that file is
    returned if it exists). ``ftw``: ftw-tools' checkpoint URL, which
    Ultralytics resolves relative to the current working directory, is
    downloaded into the current working directory, so the later run must
    start from the same directory to find it offline.

    Returns
    -------
    list[str]
        Local paths of the weights.
    """
    opts = DAOptions.from_engine_params(config.engine_params)
    if opts.checkpoint_path:
        path = Path(opts.checkpoint_path).expanduser()
        if not path.is_file():
            raise FileNotFoundError(f"checkpoint_path does not exist: {path}")
        return [str(path.resolve())]
    if opts.backend in ("native", "reference"):
        return [download_da_weights(opts.model_key)]
    url = _ftw_checkpoint_url(opts.model.ftw_name)
    if url is None:
        raise RuntimeError(
            f"The installed ftw-tools has no checkpoint URL for {opts.model.ftw_name!r}."
        )
    from urllib.parse import unquote, urlparse

    import torch

    target = Path.cwd() / Path(unquote(urlparse(url).path)).name
    if not target.is_file():
        torch.hub.download_url_to_file(url, str(target), progress=True)
    return [str(target)]

Fields of The World

ftw

FTW (Fields of The World) semantic-segmentation engine.

Runs an ftw-tools checkpoint on R, G, B and NIR and polygonises the predicted field class (1) with ftw_tools.postprocess.polygonize.polygonize. The default model is the ftw-tools MODEL_REGISTRY entry marked default (FTW_PRUE_EFNET_B5 in ftw-tools 2.0.0b5: a PRUE U-Net with an EfficientNet-B5 encoder, two input windows). Other registry models are chosen with engine_params["model"] (see :func:list_ftw_models); a local checkpoint with engine_params["checkpoint_path"]. Instance-segmentation entries of the registry (Delineate-Anything) are rejected; use the delineate-anything engine for those.

Input windows

The number of windows follows the model: registry models use ModelSpec.requires_window; for a checkpoint file in_channels is read from its hyper_parameters (4 = one window, 8 = two windows).

Two-window models take [R, G, B, NIR] of an early-season window A followed by the same bands of a late-season window B, the band order that ftw-tools' own inference input builder writes (ftw inference download: create_input(win_a, win_b) stacks the scenes time-major, B04, B03, B02, B08 per scene) and that ftw_tools.inference.inference.run passes to the model unchanged. The window centres are FTW's summer-crop start and end of season over the study-area bounding box (ftw_tools.utils.get_harvest_integer_from_bbox and harvest_to_datetime), with the end of season placed in year + 1 when it falls before the start (southern-hemisphere seasons), as ftw_tools.download.download_img.scene_selection does. Each window is a median composite over centre +/- window_days (engine parameter, default 30) built by the source's composite builder with date_range set, so each window has its own cache entry. engine_params["window_dates"] (two "YYYY-MM-DD" centres) replaces the crop calendar. If a window has no imagery (the composite builder raises :class:~agribound.composites.base.NoDataError) the run fails with a :class:RuntimeError (the :class:NoDataError chained as its cause), unless engine_params["allow_annual_fallback"] is True, in which case the annual composite is used for that window (logged as a WARNING and recorded in engine_meta); other builder errors (invalid configuration, authentication, quota, network) always propagate. Each window's record in engine_meta["windows"] also holds its composite's image count (n_images), valid-pixel fraction (valid_fraction) and cloud mask (cloud_mask), read from the composite's tags. A short window has few images, and the default Sentinel-2 mask (SCL classes 3, 8, 9 and 10) can leave haze or thin cloud in the median, which the valid fraction does not reveal (a Namoi test window B of 11 images had haze over part of the study area with valid_fraction 1.0). Look at the window composites (engine_meta["windows"][...]["raster"]) or try s2_cloud_mask="cloud_score_plus". :meth:FTWEngine.stage_inputs builds these inputs without running inference (used by :mod:agribound.hpc.tiles to stage them on a node with network access). For source="local" a two-window model needs either stacked_windows=True with a raster whose bands 1-4 and 5-8 are the two windows (R, G, B, NIR each), or allow_annual_fallback=True (the single raster is used for both windows).

Single-window models use the input raster.

Radiometry

ftw-tools' default preprocessing divides the input by 3000, i.e. it expects Sentinel-2 L2A surface reflectance x 10000. Sentinel-2, Landsat and HLS composites are already on that scale after the 1.0 harmonisation and are used unchanged (Landsat and HLS are nevertheless outside the Sentinel-2 training distribution: a WARNING is logged and engine_meta["out_of_distribution_source"] is True); other value scales are converted with :func:agribound.io.raster.to_s2_dn. local rasters need engine_params["value_scale"]. NaN, infinite and declared nodata values are replaced with 0 before the input raster is written (:func:write_ftw_input).

Polygonisation

polygonize applies simplify and the morphology options in the units of the prediction raster's CRS, and computes areas (min_size) in those units when they are metres. A prediction raster in a geographic CRS, a CRS whose linear unit is not the metre (e.g. US survey feet) or a Mercator/Web Mercator CRS is therefore reprojected (nearest neighbour) to the UTM zone of the study-area centre first (ftw-baselines issue #271; :func:metric_reprojection_reason), so simplify and min_size are in metres and m² of UTM or of the raster's own metric projection (other metric projections, e.g. Albers, keep their small scale distortion). polygonize processes the mask in windows of polygonization_stride pixels (default 2048); fields crossing a window edge are split there unless merge_adjacent is set. close_interiors (default True) fills all holes; on ftw-tools builds whose polygonize cannot close the interiors of a MultiPolygon (2.0.0b5), combining it with erode_dilate or dilate_erode raises :class:ValueError before any input is built.

Engine parameters (all optional)

model, checkpoint_path, window_days (30), window_dates, allow_annual_fallback (False), stacked_windows (False), value_scale (local rasters), resize_factor (2), patch_size (default: :func:select_patch_size), batch_size (2), padding (ftw-tools default), softmax_threshold (polygonise field probability >= threshold instead of the arg-max class; implies save_scores), save_scores (only together with softmax_threshold), simplify (metres, 0; the pipeline simplifies later with config.simplify_tolerance), close_interiors (True), merge_adjacent (None), polygonization_stride (2048), max_size (m², None), erode_dilate, dilate_erode, erode_dilate_raster, dilate_erode_raster (0) and thin_boundaries (False).

gdf.attrs["engine_meta"] records the backend ("ftw-tools"), the ftw-tools version, the model key, the checkpoint URL, the path and SHA-256 of the checkpoint file (for registry models, the copy ftw-tools caches under torch.hub.get_dir()/checkpoints), the window dates and how they were chosen, and the input units.

References

Kerner, H., et al. (2025). Fields of The World: A Machine Learning Benchmark Dataset for Global Agricultural Field Boundary Segmentation. Proceedings of the AAAI Conference on Artificial Intelligence 39(27), 28151-28159. doi:10.1609/aaai.v39i27.35034.

Muhawenayo, G., et al. (2026). PRUE: A Practical Recipe for Field Boundary Segmentation at Scale. arXiv:2603.27101.

FTWEngine

Bases: DelineationEngine

Field boundary delineation with FTW semantic segmentation (see the module docstring).

Source code in agribound/engines/ftw.py
 785
 786
 787
 788
 789
 790
 791
 792
 793
 794
 795
 796
 797
 798
 799
 800
 801
 802
 803
 804
 805
 806
 807
 808
 809
 810
 811
 812
 813
 814
 815
 816
 817
 818
 819
 820
 821
 822
 823
 824
 825
 826
 827
 828
 829
 830
 831
 832
 833
 834
 835
 836
 837
 838
 839
 840
 841
 842
 843
 844
 845
 846
 847
 848
 849
 850
 851
 852
 853
 854
 855
 856
 857
 858
 859
 860
 861
 862
 863
 864
 865
 866
 867
 868
 869
 870
 871
 872
 873
 874
 875
 876
 877
 878
 879
 880
 881
 882
 883
 884
 885
 886
 887
 888
 889
 890
 891
 892
 893
 894
 895
 896
 897
 898
 899
 900
 901
 902
 903
 904
 905
 906
 907
 908
 909
 910
 911
 912
 913
 914
 915
 916
 917
 918
 919
 920
 921
 922
 923
 924
 925
 926
 927
 928
 929
 930
 931
 932
 933
 934
 935
 936
 937
 938
 939
 940
 941
 942
 943
 944
 945
 946
 947
 948
 949
 950
 951
 952
 953
 954
 955
 956
 957
 958
 959
 960
 961
 962
 963
 964
 965
 966
 967
 968
 969
 970
 971
 972
 973
 974
 975
 976
 977
 978
 979
 980
 981
 982
 983
 984
 985
 986
 987
 988
 989
 990
 991
 992
 993
 994
 995
 996
 997
 998
 999
1000
1001
1002
1003
1004
1005
1006
1007
1008
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058
1059
1060
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
class FTWEngine(DelineationEngine):
    """Field boundary delineation with FTW semantic segmentation (see the module docstring)."""

    name = "ftw"
    supported_sources = list(_REGISTRY_ENTRY["supported_sources"])
    requires_bands = list(_REGISTRY_ENTRY["requires_bands"])

    def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
        """Run FTW inference and polygonisation.

        Parameters
        ----------
        raster_path : str
            Annual composite (or local raster) for the run.
        config : AgriboundConfig
            Pipeline configuration (``engine_params``: module docstring).

        Returns
        -------
        geopandas.GeoDataFrame
            Field polygons (in a projected CRS) with
            ``gdf.attrs["engine_meta"]``.
        """
        try:
            from ftw_tools.inference.inference import run as ftw_run
            from ftw_tools.postprocess.polygonize import polygonize as ftw_polygonize
        except ImportError:
            raise ImportError(
                "ftw-tools (>= 2.0.0b5) is required for the FTW engine. Install with: "
                "pip install 'agribound[ftw]'"
            ) from None
        from agribound._cache import cache_path
        from agribound.registry import source_value_scale

        self.validate_input(raster_path, config)
        params = dict(config.engine_params)
        choice = resolve_ftw_model(params)
        value_scale = params.get("value_scale") or source_value_scale(config.source)
        if value_scale in ("dn", "unknown"):
            raise ValueError(
                f"FTW needs Sentinel-2-like reflectance, but source {config.source!r} has value "
                f"scale {value_scale!r}. For a local raster set engine_params['value_scale'] to "
                "'reflectance_x10000', 'unit' (0-1 reflectance) or 'uint8'."
            )
        softmax_threshold = params.get("softmax_threshold")
        save_scores = bool(params.get("save_scores", softmax_threshold is not None))
        if softmax_threshold is not None:
            softmax_threshold = float(softmax_threshold)
            if not 0 < softmax_threshold < 1:
                raise ValueError(f"softmax_threshold must be in (0, 1), got {softmax_threshold}")
            save_scores = True
        elif save_scores:
            raise ValueError(
                "save_scores=True writes class probabilities, which ftw-tools polygonize only "
                "reads with a softmax_threshold; set engine_params['softmax_threshold']."
            )

        # Polygonize options are checked before any (GEE) input is built.
        poly_kwargs: dict[str, Any] = {
            "simplify": float(params.get("simplify", 0)),
            "min_size": float(config.min_field_area_m2),
            "close_interiors": bool(params.get("close_interiors", True)),
        }
        for key in _POLYGONIZE_PARAMS:
            if key in params and key != "close_interiors":
                poly_kwargs[key] = params[key]
        if softmax_threshold is not None:
            poly_kwargs["softmax_threshold"] = softmax_threshold
        accepted = inspect.signature(ftw_polygonize).parameters
        unknown = sorted(k for k in poly_kwargs if k not in accepted)
        if unknown:
            raise ValueError(f"The installed ftw-tools polygonize() does not accept {unknown}")
        _check_close_interiors(ftw_polygonize, poly_kwargs)

        rgbn = get_canonical_band_indices(config.source, _RGBN, bands=config.bands)
        out_of_distribution = config.source in _OUT_OF_DISTRIBUTION_SOURCES
        if out_of_distribution:
            logger.warning(
                "FTW checkpoints are trained on Sentinel-2 L2A; %s input (harmonised surface "
                "reflectance x 10000) is out of distribution and accuracy is not established.",
                config.source,
            )
        meta: dict[str, Any] = {
            "backend": "ftw-tools",
            "ftw_tools_version": _ftw_version(),
            "model": choice.registry_key or "checkpoint",
            "checkpoint_url": choice.url,
            "checkpoint_path": choice.checkpoint_path,
            "checkpoint_sha256": choice.checkpoint_sha256,
            "model_license": choice.license,
            "model_version": choice.version,
            "in_channels": choice.in_channels,
            "n_windows": choice.n_windows,
            "band_indices_rgbn": rgbn,
            "value_scale": value_scale,
            "out_of_distribution_source": out_of_distribution,
            "input_units": (
                "S2 L2A reflectance x10000 ("
                + (
                    "composite values copied unchanged"
                    if value_scale == "reflectance_x10000"
                    else f"converted from {value_scale} with agribound.io.raster.to_s2_dn"
                )
                + "; NaN, inf and declared nodata -> 0); ftw-tools divides by 3000"
            ),
        }

        # --- build the FTW input -------------------------------------------
        sources, meta["windows"] = self._input_sources(
            config, raster_path, params, choice.n_windows, rgbn
        )

        window_key = json.dumps(
            [(_fingerprint(path), bands) for path, bands in sources], sort_keys=True
        )
        ftw_input = cache_path(
            config,
            "ftw_input",
            ".tif",
            choice.n_windows,
            window_key,
            value_scale,
            FTW_INPUT_VERSION,
        )
        if not ftw_input.exists():
            write_ftw_input(ftw_input, sources, config.source, value_scale=value_scale)
        meta["ftw_input"] = str(ftw_input)

        # --- inference ------------------------------------------------------
        import rasterio

        device = config.resolve_device()
        with rasterio.open(ftw_input) as src:
            patch_size = select_patch_size(src.height, src.width, params.get("patch_size"))
        run_kwargs: dict[str, Any] = {
            "input": str(ftw_input),
            "model": choice.run_model,
            "resize_factor": int(params.get("resize_factor", 2)),
            "gpu": 0 if device == "cuda" else -1,
            "patch_size": patch_size,
            "batch_size": int(params.get("batch_size", 2)),
            "num_workers": config.n_workers,
            "padding": params.get("padding"),
            "overwrite": True,
            "mps_mode": device == "mps",
            "save_scores": save_scores,
        }
        if "nan_fill_value" in inspect.signature(ftw_run).parameters:
            run_kwargs["nan_fill_value"] = 0.0
        pred = cache_path(
            config,
            "ftw_pred",
            ".tif",
            choice.cache_id,
            # Registry weights are identified by their URL (a checkpoint by its SHA-256
            # in cache_id), so a registry update for the same key is not reused.
            choice.url,
            _fingerprint(ftw_input),
            json.dumps(
                {k: v for k, v in run_kwargs.items() if k not in ("input", "model", "num_workers")},
                sort_keys=True,
                default=str,
            ),
            # Preprocessing and patch stitching can change between ftw-tools releases.
            meta["ftw_tools_version"],
        )
        meta.update(
            {
                "device": device,
                "resize_factor": run_kwargs["resize_factor"],
                "patch_size": run_kwargs["patch_size"],
                "batch_size": run_kwargs["batch_size"],
                "save_scores": save_scores,
            }
        )
        if pred.exists():
            logger.info("Using cached FTW prediction: %s", pred)
            meta["cached_prediction"] = True
        else:
            logger.info("Running FTW inference (model=%s, device=%s)", meta["model"], device)
            # Written under a temporary name and renamed, so an interrupted run
            # never leaves a partial raster that a later run would take as cached.
            partial = pred.with_name(pred.stem + ".partial.tif")
            ftw_run(out=str(partial), **run_kwargs)
            if not partial.exists():
                raise RuntimeError(f"FTW inference did not write the prediction raster {partial}")
            os.replace(partial, pred)
        if choice.registry_key is not None:
            # Pin the exact registry weights: ftw-tools loads the cached file.
            meta.update(registry_checkpoint_facts(choice.registry_key))

        # --- polygonise in metres -------------------------------------------
        poly_input = str(pred)
        with rasterio.open(pred) as src:
            pred_crs = src.crs
        reason = metric_reprojection_reason(pred_crs)
        if reason is not None:
            from shapely.geometry import box

            from agribound.io.crs import reproject_raster, utm_crs_for_geometry

            utm = utm_crs_for_geometry(box(*_aoi_bounds_4326(config, raster_path)))
            target = pred.with_name(pred.stem + f"_epsg{utm.to_epsg()}.tif")
            if not target.exists():
                partial = target.with_name(target.stem + ".partial.tif")
                reproject_raster(pred, partial, utm, resampling="nearest")
                os.replace(partial, target)
            poly_input = str(target)
            meta["prediction_reprojected_to"] = f"EPSG:{utm.to_epsg()}"
            meta["prediction_reprojection_reason"] = reason
        meta["prediction_crs"] = str(pred_crs)
        meta["polygonize"] = {
            ("simplify_m" if key == "simplify" else key): value
            for key, value in poly_kwargs.items()
        }

        if not _has_field_pixels(poly_input, softmax_threshold):
            logger.warning("FTW predicted no field pixels")
            with rasterio.open(poly_input) as src:
                crs = src.crs
            gdf = gpd.GeoDataFrame({"geometry": []}, geometry="geometry", crs=crs)
            meta["n_output"] = 0
            gdf.attrs["engine_meta"] = meta
            return gdf

        poly_path = cache_path(
            config, "ftw_polygons", ".gpkg", _fingerprint(poly_input), sorted(poly_kwargs.items())
        )
        ftw_polygonize(input=poly_input, out=str(poly_path), overwrite=True, **poly_kwargs)
        if not poly_path.exists():
            raise RuntimeError(f"FTW polygonization did not write {poly_path}")
        gdf = gpd.read_file(poly_path)
        meta["n_output"] = len(gdf)
        logger.info("FTW delineated %d field polygons", len(gdf))
        gdf.attrs["engine_meta"] = meta
        return gdf

    @staticmethod
    def stage_inputs(config: AgriboundConfig, raster_path: str) -> dict[str, Any]:
        """Build (or reuse from the cache) the input rasters :meth:`delineate` reads.

        Runs the input stage of :meth:`delineate` for *config* and
        *raster_path* without running inference: the model is resolved from
        ``engine_params`` with :func:`resolve_ftw_model`, as in
        :meth:`delineate` (ftw-tools' model registry for registry models, the
        checkpoint's channel count for ``checkpoint_path``), and for a
        two-window model on an Earth Engine
        source both seasonal window composites are built by the source's
        composite builder with the same window centres, ``window_days`` and
        cache keys, so a later :meth:`delineate` with the same configuration
        and cache directory finds them without network access (the crop
        calendar comes from ftw-tools' cache; see :meth:`prefetch`).
        Single-window models and local sources build nothing.

        Parameters
        ----------
        config : AgriboundConfig
            Pipeline configuration.
        raster_path : str
            Annual composite (or local raster) of the run.

        Returns
        -------
        dict
            ``n_windows`` (1 or 2), ``band_indices_rgbn`` (1-based R, G, B, NIR
            indices read from each raster), ``windows`` (the record stored in
            ``engine_meta["windows"]``: ``"single"``, or ``"a"``/``"b"`` with
            ``start``, ``end``, ``raster`` and ``status``, plus the window
            ``centres`` for Earth Engine sources) and ``rasters`` (the raster
            of each window in input order, A then B).

        Raises
        ------
        RuntimeError
            If a window has no imagery and ``allow_annual_fallback`` is not
            set (the builder's :class:`~agribound.composites.base.NoDataError`
            is its ``__cause__``).
        ValueError
            For invalid window parameters or band mappings.
        """
        params = dict(config.engine_params)
        choice = resolve_ftw_model(params)
        rgbn = get_canonical_band_indices(config.source, _RGBN, bands=config.bands)
        sources, windows = FTWEngine._input_sources(
            config, raster_path, params, choice.n_windows, rgbn
        )
        return {
            "n_windows": int(choice.n_windows),
            "band_indices_rgbn": list(rgbn),
            "windows": windows,
            "rasters": [str(path) for path, _bands in sources],
        }

    @staticmethod
    def _input_sources(
        config: AgriboundConfig,
        raster_path: str,
        params: dict[str, Any],
        n_windows: int,
        rgbn: list[int],
    ) -> tuple[list[tuple[str, list[int]]], dict[str, Any]]:
        """``(sources, windows record)`` of the FTW input for *n_windows*."""
        if n_windows == 1:
            record = {"single": {"raster": raster_path, "status": "input raster"}}
            return [(raster_path, rgbn)], record
        if config.is_gee_source():
            return FTWEngine._two_windows(config, raster_path, params, rgbn)
        return FTWEngine._two_windows_local(raster_path, params, rgbn)

    @staticmethod
    def _two_windows(
        config: AgriboundConfig, raster_path: str, params: dict[str, Any], rgbn: list[int]
    ) -> tuple[list[tuple[str, list[int]]], dict[str, Any]]:
        days = int(params.get("window_days", 30))
        if days < 1:
            raise ValueError(f"window_days must be >= 1, got {days}")
        allow_fallback = bool(params.get("allow_annual_fallback", False))
        centres = _window_centres(config, raster_path, params)
        record: dict[str, Any] = {"centres": centres, "window_days": days}
        sources = []
        for label in ("a", "b"):
            start, end = window_range(centres[f"centre_{label}"], days)
            path, status, error = _build_window(
                config, label.upper(), start, end, raster_path, allow_fallback
            )
            record[label] = {"start": start, "end": end, "raster": path, "status": status}
            record[label].update(_window_composite_facts(path))
            if error:
                record[label]["error"] = error
            sources.append((path, rgbn))
        logger.info(
            "FTW windows: A %s..%s (%s), B %s..%s (%s)",
            record["a"]["start"],
            record["a"]["end"],
            record["a"]["status"],
            record["b"]["start"],
            record["b"]["end"],
            record["b"]["status"],
        )
        return sources, record

    @staticmethod
    def _two_windows_local(
        raster_path: str, params: dict[str, Any], rgbn: list[int]
    ) -> tuple[list[tuple[str, list[int]]], dict[str, Any]]:
        import rasterio

        if params.get("stacked_windows"):
            with rasterio.open(raster_path) as src:
                if src.count < 8:
                    raise ValueError(
                        f"stacked_windows=True needs at least 8 bands, {raster_path} has "
                        f"{src.count}"
                    )
            return [(raster_path, [1, 2, 3, 4]), (raster_path, [5, 6, 7, 8])], {
                "a": {"raster": raster_path, "bands": [1, 2, 3, 4], "status": "stacked"},
                "b": {"raster": raster_path, "bands": [5, 6, 7, 8], "status": "stacked"},
            }
        if params.get("allow_annual_fallback"):
            logger.warning(
                "Local raster %s is used for both FTW windows (allow_annual_fallback=True); the "
                "input is not seasonal.",
                raster_path,
            )
            return [(raster_path, rgbn), (raster_path, rgbn)], {
                "a": {"raster": raster_path, "status": "annual_fallback"},
                "b": {"raster": raster_path, "status": "annual_fallback"},
            }
        raise ValueError(
            "This FTW model needs two seasonal windows, which cannot be built from a local "
            "raster. Provide an 8-band raster (window A R, G, B, NIR; window B R, G, B, NIR) "
            "with engine_params['stacked_windows']=True, use a single-window model "
            "(e.g. model='FTW_v2_3_Class_FULL_singleWindow'), or set "
            "engine_params['allow_annual_fallback']=True to use the raster for both windows."
        )

    @classmethod
    def prefetch(cls, config: AgriboundConfig) -> list[str]:
        """Download the FTW checkpoint (and the crop calendar for two-window models).

        Registry checkpoints are saved where ftw-tools' ``run()`` looks for
        them (``torch.hub.get_dir()/checkpoints/<model>.ckpt``); the crop
        calendar goes to ``$FTW_CACHE_DIR/crop_calendar`` (default
        ``~/.cache/ftw-tools``). Set ``TORCH_HOME`` and ``FTW_CACHE_DIR`` to
        shared storage on HPC systems.

        Returns
        -------
        list[str]
            Local paths of the checkpoint and crop-calendar files.
        """
        choice = resolve_ftw_model(config.engine_params)
        paths: list[str] = []
        if choice.checkpoint_path:
            paths.append(choice.checkpoint_path)
        else:
            import torch

            target = Path(torch.hub.get_dir()) / "checkpoints" / f"{choice.registry_key}.ckpt"
            target.parent.mkdir(parents=True, exist_ok=True)
            if not target.exists():
                torch.hub.download_url_to_file(choice.url, str(target), progress=True)
            paths.append(str(target))
        if choice.n_windows == 2 and config.engine_params.get("window_dates") is None:
            from ftw_tools.download.crop_calendar import ensure_crop_calendar_exists
            from ftw_tools.settings import CROP_CAL_SUMMER_END, CROP_CAL_SUMMER_START

            calendar_dir = ensure_crop_calendar_exists()
            paths += [
                str(calendar_dir / CROP_CAL_SUMMER_START),
                str(calendar_dir / CROP_CAL_SUMMER_END),
            ]
        return paths
delineate
delineate(raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame

Run FTW inference and polygonisation.

Parameters:

Name Type Description Default
raster_path str

Annual composite (or local raster) for the run.

required
config AgriboundConfig

Pipeline configuration (engine_params: module docstring).

required

Returns:

Type Description
GeoDataFrame

Field polygons (in a projected CRS) with gdf.attrs["engine_meta"].

Source code in agribound/engines/ftw.py
def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
    """Run FTW inference and polygonisation.

    Parameters
    ----------
    raster_path : str
        Annual composite (or local raster) for the run.
    config : AgriboundConfig
        Pipeline configuration (``engine_params``: module docstring).

    Returns
    -------
    geopandas.GeoDataFrame
        Field polygons (in a projected CRS) with
        ``gdf.attrs["engine_meta"]``.
    """
    try:
        from ftw_tools.inference.inference import run as ftw_run
        from ftw_tools.postprocess.polygonize import polygonize as ftw_polygonize
    except ImportError:
        raise ImportError(
            "ftw-tools (>= 2.0.0b5) is required for the FTW engine. Install with: "
            "pip install 'agribound[ftw]'"
        ) from None
    from agribound._cache import cache_path
    from agribound.registry import source_value_scale

    self.validate_input(raster_path, config)
    params = dict(config.engine_params)
    choice = resolve_ftw_model(params)
    value_scale = params.get("value_scale") or source_value_scale(config.source)
    if value_scale in ("dn", "unknown"):
        raise ValueError(
            f"FTW needs Sentinel-2-like reflectance, but source {config.source!r} has value "
            f"scale {value_scale!r}. For a local raster set engine_params['value_scale'] to "
            "'reflectance_x10000', 'unit' (0-1 reflectance) or 'uint8'."
        )
    softmax_threshold = params.get("softmax_threshold")
    save_scores = bool(params.get("save_scores", softmax_threshold is not None))
    if softmax_threshold is not None:
        softmax_threshold = float(softmax_threshold)
        if not 0 < softmax_threshold < 1:
            raise ValueError(f"softmax_threshold must be in (0, 1), got {softmax_threshold}")
        save_scores = True
    elif save_scores:
        raise ValueError(
            "save_scores=True writes class probabilities, which ftw-tools polygonize only "
            "reads with a softmax_threshold; set engine_params['softmax_threshold']."
        )

    # Polygonize options are checked before any (GEE) input is built.
    poly_kwargs: dict[str, Any] = {
        "simplify": float(params.get("simplify", 0)),
        "min_size": float(config.min_field_area_m2),
        "close_interiors": bool(params.get("close_interiors", True)),
    }
    for key in _POLYGONIZE_PARAMS:
        if key in params and key != "close_interiors":
            poly_kwargs[key] = params[key]
    if softmax_threshold is not None:
        poly_kwargs["softmax_threshold"] = softmax_threshold
    accepted = inspect.signature(ftw_polygonize).parameters
    unknown = sorted(k for k in poly_kwargs if k not in accepted)
    if unknown:
        raise ValueError(f"The installed ftw-tools polygonize() does not accept {unknown}")
    _check_close_interiors(ftw_polygonize, poly_kwargs)

    rgbn = get_canonical_band_indices(config.source, _RGBN, bands=config.bands)
    out_of_distribution = config.source in _OUT_OF_DISTRIBUTION_SOURCES
    if out_of_distribution:
        logger.warning(
            "FTW checkpoints are trained on Sentinel-2 L2A; %s input (harmonised surface "
            "reflectance x 10000) is out of distribution and accuracy is not established.",
            config.source,
        )
    meta: dict[str, Any] = {
        "backend": "ftw-tools",
        "ftw_tools_version": _ftw_version(),
        "model": choice.registry_key or "checkpoint",
        "checkpoint_url": choice.url,
        "checkpoint_path": choice.checkpoint_path,
        "checkpoint_sha256": choice.checkpoint_sha256,
        "model_license": choice.license,
        "model_version": choice.version,
        "in_channels": choice.in_channels,
        "n_windows": choice.n_windows,
        "band_indices_rgbn": rgbn,
        "value_scale": value_scale,
        "out_of_distribution_source": out_of_distribution,
        "input_units": (
            "S2 L2A reflectance x10000 ("
            + (
                "composite values copied unchanged"
                if value_scale == "reflectance_x10000"
                else f"converted from {value_scale} with agribound.io.raster.to_s2_dn"
            )
            + "; NaN, inf and declared nodata -> 0); ftw-tools divides by 3000"
        ),
    }

    # --- build the FTW input -------------------------------------------
    sources, meta["windows"] = self._input_sources(
        config, raster_path, params, choice.n_windows, rgbn
    )

    window_key = json.dumps(
        [(_fingerprint(path), bands) for path, bands in sources], sort_keys=True
    )
    ftw_input = cache_path(
        config,
        "ftw_input",
        ".tif",
        choice.n_windows,
        window_key,
        value_scale,
        FTW_INPUT_VERSION,
    )
    if not ftw_input.exists():
        write_ftw_input(ftw_input, sources, config.source, value_scale=value_scale)
    meta["ftw_input"] = str(ftw_input)

    # --- inference ------------------------------------------------------
    import rasterio

    device = config.resolve_device()
    with rasterio.open(ftw_input) as src:
        patch_size = select_patch_size(src.height, src.width, params.get("patch_size"))
    run_kwargs: dict[str, Any] = {
        "input": str(ftw_input),
        "model": choice.run_model,
        "resize_factor": int(params.get("resize_factor", 2)),
        "gpu": 0 if device == "cuda" else -1,
        "patch_size": patch_size,
        "batch_size": int(params.get("batch_size", 2)),
        "num_workers": config.n_workers,
        "padding": params.get("padding"),
        "overwrite": True,
        "mps_mode": device == "mps",
        "save_scores": save_scores,
    }
    if "nan_fill_value" in inspect.signature(ftw_run).parameters:
        run_kwargs["nan_fill_value"] = 0.0
    pred = cache_path(
        config,
        "ftw_pred",
        ".tif",
        choice.cache_id,
        # Registry weights are identified by their URL (a checkpoint by its SHA-256
        # in cache_id), so a registry update for the same key is not reused.
        choice.url,
        _fingerprint(ftw_input),
        json.dumps(
            {k: v for k, v in run_kwargs.items() if k not in ("input", "model", "num_workers")},
            sort_keys=True,
            default=str,
        ),
        # Preprocessing and patch stitching can change between ftw-tools releases.
        meta["ftw_tools_version"],
    )
    meta.update(
        {
            "device": device,
            "resize_factor": run_kwargs["resize_factor"],
            "patch_size": run_kwargs["patch_size"],
            "batch_size": run_kwargs["batch_size"],
            "save_scores": save_scores,
        }
    )
    if pred.exists():
        logger.info("Using cached FTW prediction: %s", pred)
        meta["cached_prediction"] = True
    else:
        logger.info("Running FTW inference (model=%s, device=%s)", meta["model"], device)
        # Written under a temporary name and renamed, so an interrupted run
        # never leaves a partial raster that a later run would take as cached.
        partial = pred.with_name(pred.stem + ".partial.tif")
        ftw_run(out=str(partial), **run_kwargs)
        if not partial.exists():
            raise RuntimeError(f"FTW inference did not write the prediction raster {partial}")
        os.replace(partial, pred)
    if choice.registry_key is not None:
        # Pin the exact registry weights: ftw-tools loads the cached file.
        meta.update(registry_checkpoint_facts(choice.registry_key))

    # --- polygonise in metres -------------------------------------------
    poly_input = str(pred)
    with rasterio.open(pred) as src:
        pred_crs = src.crs
    reason = metric_reprojection_reason(pred_crs)
    if reason is not None:
        from shapely.geometry import box

        from agribound.io.crs import reproject_raster, utm_crs_for_geometry

        utm = utm_crs_for_geometry(box(*_aoi_bounds_4326(config, raster_path)))
        target = pred.with_name(pred.stem + f"_epsg{utm.to_epsg()}.tif")
        if not target.exists():
            partial = target.with_name(target.stem + ".partial.tif")
            reproject_raster(pred, partial, utm, resampling="nearest")
            os.replace(partial, target)
        poly_input = str(target)
        meta["prediction_reprojected_to"] = f"EPSG:{utm.to_epsg()}"
        meta["prediction_reprojection_reason"] = reason
    meta["prediction_crs"] = str(pred_crs)
    meta["polygonize"] = {
        ("simplify_m" if key == "simplify" else key): value
        for key, value in poly_kwargs.items()
    }

    if not _has_field_pixels(poly_input, softmax_threshold):
        logger.warning("FTW predicted no field pixels")
        with rasterio.open(poly_input) as src:
            crs = src.crs
        gdf = gpd.GeoDataFrame({"geometry": []}, geometry="geometry", crs=crs)
        meta["n_output"] = 0
        gdf.attrs["engine_meta"] = meta
        return gdf

    poly_path = cache_path(
        config, "ftw_polygons", ".gpkg", _fingerprint(poly_input), sorted(poly_kwargs.items())
    )
    ftw_polygonize(input=poly_input, out=str(poly_path), overwrite=True, **poly_kwargs)
    if not poly_path.exists():
        raise RuntimeError(f"FTW polygonization did not write {poly_path}")
    gdf = gpd.read_file(poly_path)
    meta["n_output"] = len(gdf)
    logger.info("FTW delineated %d field polygons", len(gdf))
    gdf.attrs["engine_meta"] = meta
    return gdf
stage_inputs staticmethod
stage_inputs(config: AgriboundConfig, raster_path: str) -> dict[str, Any]

Build (or reuse from the cache) the input rasters :meth:delineate reads.

Runs the input stage of :meth:delineate for config and raster_path without running inference: the model is resolved from engine_params with :func:resolve_ftw_model, as in :meth:delineate (ftw-tools' model registry for registry models, the checkpoint's channel count for checkpoint_path), and for a two-window model on an Earth Engine source both seasonal window composites are built by the source's composite builder with the same window centres, window_days and cache keys, so a later :meth:delineate with the same configuration and cache directory finds them without network access (the crop calendar comes from ftw-tools' cache; see :meth:prefetch). Single-window models and local sources build nothing.

Parameters:

Name Type Description Default
config AgriboundConfig

Pipeline configuration.

required
raster_path str

Annual composite (or local raster) of the run.

required

Returns:

Type Description
dict

n_windows (1 or 2), band_indices_rgbn (1-based R, G, B, NIR indices read from each raster), windows (the record stored in engine_meta["windows"]: "single", or "a"/"b" with start, end, raster and status, plus the window centres for Earth Engine sources) and rasters (the raster of each window in input order, A then B).

Raises:

Type Description
RuntimeError

If a window has no imagery and allow_annual_fallback is not set (the builder's :class:~agribound.composites.base.NoDataError is its __cause__).

ValueError

For invalid window parameters or band mappings.

Source code in agribound/engines/ftw.py
@staticmethod
def stage_inputs(config: AgriboundConfig, raster_path: str) -> dict[str, Any]:
    """Build (or reuse from the cache) the input rasters :meth:`delineate` reads.

    Runs the input stage of :meth:`delineate` for *config* and
    *raster_path* without running inference: the model is resolved from
    ``engine_params`` with :func:`resolve_ftw_model`, as in
    :meth:`delineate` (ftw-tools' model registry for registry models, the
    checkpoint's channel count for ``checkpoint_path``), and for a
    two-window model on an Earth Engine
    source both seasonal window composites are built by the source's
    composite builder with the same window centres, ``window_days`` and
    cache keys, so a later :meth:`delineate` with the same configuration
    and cache directory finds them without network access (the crop
    calendar comes from ftw-tools' cache; see :meth:`prefetch`).
    Single-window models and local sources build nothing.

    Parameters
    ----------
    config : AgriboundConfig
        Pipeline configuration.
    raster_path : str
        Annual composite (or local raster) of the run.

    Returns
    -------
    dict
        ``n_windows`` (1 or 2), ``band_indices_rgbn`` (1-based R, G, B, NIR
        indices read from each raster), ``windows`` (the record stored in
        ``engine_meta["windows"]``: ``"single"``, or ``"a"``/``"b"`` with
        ``start``, ``end``, ``raster`` and ``status``, plus the window
        ``centres`` for Earth Engine sources) and ``rasters`` (the raster
        of each window in input order, A then B).

    Raises
    ------
    RuntimeError
        If a window has no imagery and ``allow_annual_fallback`` is not
        set (the builder's :class:`~agribound.composites.base.NoDataError`
        is its ``__cause__``).
    ValueError
        For invalid window parameters or band mappings.
    """
    params = dict(config.engine_params)
    choice = resolve_ftw_model(params)
    rgbn = get_canonical_band_indices(config.source, _RGBN, bands=config.bands)
    sources, windows = FTWEngine._input_sources(
        config, raster_path, params, choice.n_windows, rgbn
    )
    return {
        "n_windows": int(choice.n_windows),
        "band_indices_rgbn": list(rgbn),
        "windows": windows,
        "rasters": [str(path) for path, _bands in sources],
    }
prefetch classmethod
prefetch(config: AgriboundConfig) -> list[str]

Download the FTW checkpoint (and the crop calendar for two-window models).

Registry checkpoints are saved where ftw-tools' run() looks for them (torch.hub.get_dir()/checkpoints/<model>.ckpt); the crop calendar goes to $FTW_CACHE_DIR/crop_calendar (default ~/.cache/ftw-tools). Set TORCH_HOME and FTW_CACHE_DIR to shared storage on HPC systems.

Returns:

Type Description
list[str]

Local paths of the checkpoint and crop-calendar files.

Source code in agribound/engines/ftw.py
@classmethod
def prefetch(cls, config: AgriboundConfig) -> list[str]:
    """Download the FTW checkpoint (and the crop calendar for two-window models).

    Registry checkpoints are saved where ftw-tools' ``run()`` looks for
    them (``torch.hub.get_dir()/checkpoints/<model>.ckpt``); the crop
    calendar goes to ``$FTW_CACHE_DIR/crop_calendar`` (default
    ``~/.cache/ftw-tools``). Set ``TORCH_HOME`` and ``FTW_CACHE_DIR`` to
    shared storage on HPC systems.

    Returns
    -------
    list[str]
        Local paths of the checkpoint and crop-calendar files.
    """
    choice = resolve_ftw_model(config.engine_params)
    paths: list[str] = []
    if choice.checkpoint_path:
        paths.append(choice.checkpoint_path)
    else:
        import torch

        target = Path(torch.hub.get_dir()) / "checkpoints" / f"{choice.registry_key}.ckpt"
        target.parent.mkdir(parents=True, exist_ok=True)
        if not target.exists():
            torch.hub.download_url_to_file(choice.url, str(target), progress=True)
        paths.append(str(target))
    if choice.n_windows == 2 and config.engine_params.get("window_dates") is None:
        from ftw_tools.download.crop_calendar import ensure_crop_calendar_exists
        from ftw_tools.settings import CROP_CAL_SUMMER_END, CROP_CAL_SUMMER_START

        calendar_dir = ensure_crop_calendar_exists()
        paths += [
            str(calendar_dir / CROP_CAL_SUMMER_START),
            str(calendar_dir / CROP_CAL_SUMMER_END),
        ]
    return paths

list_ftw_models

list_ftw_models(include_legacy: bool = False) -> dict[str, dict]

List the models in the installed ftw-tools MODEL_REGISTRY.

Parameters:

Name Type Description Default
include_legacy bool

Include models marked legacy (FTW v1/v2 checkpoints; default False).

False

Returns:

Type Description
dict[str, dict]

Model name -> {"url", "title", "description", "license", "version", "requires_window", "requires_polygonize", "instance_segmentation", "default", "legacy"}. Instance-segmentation entries (Delineate-Anything) are listed but are run by the delineate-anything engine, not by this one.

Raises:

Type Description
ImportError

If ftw-tools is not installed.

Examples:

>>> from agribound.engines.ftw import list_ftw_models
>>> for name, info in list_ftw_models().items():
...     print(f"{name}: {info['title']}")
Source code in agribound/engines/ftw.py
def list_ftw_models(include_legacy: bool = False) -> dict[str, dict]:
    """List the models in the installed ftw-tools ``MODEL_REGISTRY``.

    Parameters
    ----------
    include_legacy : bool
        Include models marked legacy (FTW v1/v2 checkpoints; default False).

    Returns
    -------
    dict[str, dict]
        Model name -> ``{"url", "title", "description", "license", "version",
        "requires_window", "requires_polygonize", "instance_segmentation",
        "default", "legacy"}``. Instance-segmentation entries
        (Delineate-Anything) are listed but are run by the
        ``delineate-anything`` engine, not by this one.

    Raises
    ------
    ImportError
        If ftw-tools is not installed.

    Examples
    --------
    >>> from agribound.engines.ftw import list_ftw_models
    >>> for name, info in list_ftw_models().items():
    ...     print(f"{name}: {info['title']}")
    """
    models = {}
    for name, spec in _model_registry().items():
        if not include_legacy and spec.legacy:
            continue
        models[name] = {
            "url": spec.url,
            "title": spec.title,
            "description": spec.description,
            "license": spec.license,
            "version": spec.version,
            "requires_window": spec.requires_window,
            "requires_polygonize": spec.requires_polygonize,
            "instance_segmentation": spec.instance_segmentation,
            "default": spec.default,
            "legacy": spec.legacy,
        }
    return models

GeoAI

geoai_field

GeoAI Mask R-CNN field instance segmentation (geoai-py).

Uses geoai's instance-segmentation workflow (Wu, 2026, JOSS 11(118):9605): geoai.train.instance_segmentation runs a torchvision Mask R-CNN ResNet50-FPN (2 classes: background, field) with a sliding window and class-aware NMS and writes an instance-id raster; each instance is then vectorised with geoai.utils.raster.raster_to_vector.

No field-boundary weights are published for geoai: as of 2026-09 the Hugging Face repository giswqs/geoai holds building, car, ship, solar-panel, parking-spot, water and wetland models and DINOv3 backbone weights; the field_boundary_detector.pth that geoai.AgricultureFieldDelineator names by default is not among them, and geoai's default detector weights (building_footprints_usa.pth) detect buildings. The engine therefore needs a checkpoint: from fine_tune=True with reference boundaries (agribound.engines.finetune._geoai) or given as engine_params["checkpoint_path"] -- a Mask R-CNN ResNet50-FPN state dict with 2 classes and 3 input channels, as written by geoai's train_MaskRCNN_model/agribound fine-tuning. It never falls back to other weights.

Input: canonical R, G, B bands with a scene-level 1-99 percentile stretch to uint8 (:func:agribound.engines.finetune._data.write_rgb_input), the same radiometry as the fine-tuning chips; geoai divides by 255. For source="local" without config.bands bands 1, 2, 3 are read as R, G, B.

Scale: torchvision's Mask R-CNN resizes every input image so that its shorter side is 800 px (GeneralizedRCNNTransform, min_size=800) at training and at inference. The apparent size of a field therefore depends on the image size: a 256 px training chip is enlarged 3.125 times, a 512 px inference window 1.5625 times. The inference window defaults to the training chip size recorded next to the checkpoint so that fields appear at the scale the model was trained on (:func:plan_geoai_windows).

GeoAIEngine

Bases: DelineationEngine

Field delineation with a fine-tuned geoai Mask R-CNN (see module docstring).

Source code in agribound/engines/geoai_field.py
class GeoAIEngine(DelineationEngine):
    """Field delineation with a fine-tuned geoai Mask R-CNN (see module docstring)."""

    name = "geoai"
    supported_sources = list(ENGINE_REGISTRY["geoai"]["supported_sources"])
    requires_bands = list(ENGINE_REGISTRY["geoai"]["requires_bands"])

    def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
        """Run geoai instance segmentation on a composite.

        Parameters
        ----------
        raster_path : str
            Composite GeoTIFF.
        config : AgriboundConfig
            Pipeline configuration. ``engine_params``:

            - ``checkpoint_path``: Mask R-CNN weights (``.pth``); see
              :func:`resolve_geoai_checkpoint` for the Hugging Face option.
            - ``window_size``: sliding window in pixels (default: the
              training chip size recorded next to the checkpoint, else
              geoai's 512). Keep it equal to the training chip size: Mask
              R-CNN resizes each window to an 800 px shorter side, so the
              window size sets the apparent field size (a different value is
              logged at WARNING). Windows that extend past a small raster are
              zero-padded by geoai. See :func:`plan_geoai_windows`.
            - ``overlap``: window overlap in pixels (default: half the
              window).
            - ``batch_size``: windows per forward pass (default 4; 1 when a
              raster side is shorter than the window, see
              :func:`plan_geoai_windows`).
            - ``confidence_threshold`` (default 0.5) and ``nms_threshold``
              (default 0.3, geoai's cross-window NMS).
            - ``merge_window_seams`` (default *True*), ``seam_min_px``
              (default 16) and ``seam_max_gap_px`` (default 2): join
              instances that a field was split into at the window edges and
              fill thin gaps along those edges (see
              :func:`merge_window_seams`).
            - ``clean_instance_mask``: *False* (default), *True* or a dict of
              keyword arguments for ``geoai.utils.raster.clean_instance_mask``
              (removes small instances, fills holes, smooths boundaries;
              runtime grows with instances x pixels).

        Returns
        -------
        geopandas.GeoDataFrame
            One polygon per detected field with ``instance_id`` and ``score``
            columns and ``gdf.attrs["engine_meta"]``.

        Notes
        -----
        geoai does not expose two limits of torchvision's Mask R-CNN
        (:func:`maskrcnn_limits`, recorded in ``engine_meta``): at most 100
        detections are kept per window (``box_detections_per_img``), so
        where more fields fit in one window (small fields at 10-30 m) the
        rest are lost, and detections scoring below 0.05
        (``box_score_thresh``) are discarded, so a ``confidence_threshold``
        below 0.05 has the same effect as 0.05 (logged at WARNING). For
        dense small fields, fine-tune with a smaller
        ``engine_params["chip_size"]``; inference then uses windows of the
        same size.
        """
        try:
            import torch  # noqa: F401
            from geoai.train import instance_segmentation
        except ImportError:
            raise ImportError(
                "geoai-py is required for the GeoAI engine. "
                "Install with: pip install agribound[geoai]"
            ) from None
        import rasterio

        from agribound._cache import cache_path
        from agribound.engines.base import get_canonical_band_indices
        from agribound.engines.finetune._data import (
            file_sha256,
            mask_invalid_predictions,
            package_versions,
            raster_fingerprint,
            read_json,
            read_training_meta,
            write_json,
            write_rgb_input,
        )

        self.validate_input(raster_path, config)
        params = config.engine_params
        checkpoint, hub_info = resolve_geoai_checkpoint(params)
        indices = get_canonical_band_indices(config.source, ["R", "G", "B"], bands=config.bands)
        confidence = float(params.get("confidence_threshold", 0.5))
        nms = float(params.get("nms_threshold", 0.3))
        limits = maskrcnn_limits()
        if confidence < limits["box_score_thresh"]:
            logger.warning(
                "GeoAI confidence_threshold %.3g is below Mask R-CNN's box_score_thresh %.3g, "
                "which discards lower-scoring detections first; it acts as %.3g",
                confidence,
                limits["box_score_thresh"],
                limits["box_score_thresh"],
            )
        training = read_training_meta(checkpoint)
        with rasterio.open(raster_path) as src:
            height, width = src.height, src.width
        plan = plan_geoai_windows(
            height,
            width,
            window_size=params.get("window_size"),
            overlap=params.get("overlap"),
            batch_size=int(params.get("batch_size", 4)),
            training_chip_size=training.get("chip_size"),
            min_size=int(limits["min_size"]),
        )
        window, overlap, batch_size = plan["window_size"], plan["overlap"], plan["batch_size"]
        clean = params.get("clean_instance_mask", False)

        device = config.resolve_device()
        if device == "mps":
            logger.warning(
                "GeoAI Mask R-CNN runs on CPU instead of MPS: torchvision Mask R-CNN on MPS "
                "reports Metal command-buffer errors and its detections differ from CPU "
                "(checked with torch 2.10 and geoai-py 0.43.1)"
            )
            device = "cpu"

        raster_fp = raster_fingerprint(raster_path)
        rgb_path = cache_path(config, "geoai_rgb", ".tif", raster_fp, f"bands={indices}", "uint8")
        rgb_info_path = rgb_path.with_suffix(".json")
        if rgb_path.exists() and rgb_info_path.exists():
            rgb_info = read_json(rgb_info_path)
        else:
            rgb_info = write_rgb_input(raster_path, rgb_path, indices, unit_float=False)
            write_json(rgb_info_path, rgb_info)

        inst_path = cache_path(
            config,
            "geoai_instances",
            ".tif",
            raster_fp,
            raster_fingerprint(checkpoint),
            f"bands={indices}",
            f"window={window}",
            f"overlap={overlap}",
            f"conf={confidence}",
            f"nms={nms}",
        )
        score_path = inst_path.with_name(f"{inst_path.stem}_score{inst_path.suffix}")
        inst_info_path = inst_path.with_suffix(".json")
        run_info: dict[str, Any] = {"device": device}
        if inst_path.exists() and score_path.exists() and inst_info_path.exists():
            logger.info("Using cached GeoAI instances: %s", inst_path)
            run_info = {**read_json(inst_info_path), "cache_reused": True}
        else:
            logger.info("Running GeoAI instance segmentation (device=%s)", device)
            instance_segmentation(
                input_path=str(rgb_path),
                output_path=str(inst_path),
                model_path=checkpoint,
                window_size=window,
                overlap=overlap,
                confidence_threshold=confidence,
                nms_threshold=nms,
                batch_size=batch_size,
                num_channels=3,
                num_classes=2,
                vectorize=False,
                device=device,
            )
            mask_invalid_predictions(inst_path, raster_path, indices)
            write_json(inst_info_path, run_info)

        vector_source = str(inst_path)
        seam_info: dict[str, Any] = {"seam_merge": False}
        if params.get("merge_window_seams", True):
            seam_min = int(params.get("seam_min_px", 16))
            seam_gap = int(params.get("seam_max_gap_px", 2))
            merged_path = inst_path.with_name(
                f"{inst_path.stem}_seams{seam_min}_gap{seam_gap}{inst_path.suffix}"
            )
            seam_json = merged_path.with_suffix(".json")
            if merged_path.exists() and seam_json.exists():
                seam_info = read_json(seam_json)
            else:
                seam_info = merge_window_seams(
                    str(inst_path),
                    str(merged_path),
                    window,
                    overlap,
                    min_seam_px=seam_min,
                    max_gap_px=seam_gap,
                )
                write_json(seam_json, seam_info)
            logger.info(
                "GeoAI: joined %d instances split at the %d px window edges",
                seam_info["n_instances_merged_at_seams"],
                window,
            )
            vector_source = str(merged_path)
        if clean:
            from geoai.utils.raster import clean_instance_mask

            kwargs = dict(clean) if isinstance(clean, dict) else {}
            # Clean the seam-merged raster when there is one, so the join is kept.
            src = Path(vector_source)
            cleaned = src.with_name(f"{src.stem}_cleaned{src.suffix}")
            vector_source = clean_instance_mask(str(src), str(cleaned), **kwargs)

        gdf = instances_to_polygons(vector_source, str(score_path))
        gdf.attrs["engine_meta"] = {
            "backend": "geoai.instance_segmentation",
            **package_versions("geoai-py", "torch", "torchvision"),
            "model": "maskrcnn_resnet50_fpn",
            "num_classes": 2,
            "checkpoint": str(Path(checkpoint).resolve()),
            "checkpoint_sha256": file_sha256(checkpoint),
            **({"hub": hub_info} if hub_info else {}),
            "training": training or None,
            "band_indices": indices,
            "input": rgb_info,
            "window_size": window,
            "overlap": overlap,
            "window_source": plan["window_source"],
            "training_chip_size": plan["training_chip_size"],
            "resize_factor": plan["resize_factor"],
            "confidence_threshold": confidence,
            "nms_threshold": nms,
            "batch_size": batch_size,
            "requested_batch_size": plan["requested_batch_size"],
            "maskrcnn_limits": limits,
            "clean_instance_mask": clean,
            **seam_info,
            **run_info,
        }
        if len(gdf) == 0:
            logger.warning("No field boundaries detected by GeoAI")
        logger.info("GeoAI delineated %d field boundaries", len(gdf))
        return gdf

    @classmethod
    def prefetch(cls, config: AgriboundConfig) -> list[str]:
        """Download the weights GeoAI inference and fine-tuning need offline.

        - torchvision's COCO Mask R-CNN ResNet50-FPN weights
          (``MaskRCNN_ResNet50_FPN_Weights.DEFAULT``): geoai builds its model
          with them before loading a checkpoint, and fine-tuning starts from
          them. They go to ``$TORCH_HOME/hub/checkpoints``.
        - The Hugging Face checkpoint, when ``engine_params["repo_id"]`` is
          set.

        Returns
        -------
        list[str]
            Local paths.
        """
        import os

        import torch
        from torchvision.models.detection import MaskRCNN_ResNet50_FPN_Weights

        weights = MaskRCNN_ResNet50_FPN_Weights.DEFAULT
        weights.get_state_dict(progress=True)
        paths = [os.path.join(torch.hub.get_dir(), "checkpoints", os.path.basename(weights.url))]
        params = config.engine_params
        if params.get("repo_id"):
            paths.append(resolve_geoai_checkpoint(params)[0])
        elif params.get("checkpoint_path") and Path(params["checkpoint_path"]).is_file():
            paths.append(str(Path(params["checkpoint_path"]).resolve()))
        return paths
delineate
delineate(raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame

Run geoai instance segmentation on a composite.

Parameters:

Name Type Description Default
raster_path str

Composite GeoTIFF.

required
config AgriboundConfig

Pipeline configuration. engine_params:

  • checkpoint_path: Mask R-CNN weights (.pth); see :func:resolve_geoai_checkpoint for the Hugging Face option.
  • window_size: sliding window in pixels (default: the training chip size recorded next to the checkpoint, else geoai's 512). Keep it equal to the training chip size: Mask R-CNN resizes each window to an 800 px shorter side, so the window size sets the apparent field size (a different value is logged at WARNING). Windows that extend past a small raster are zero-padded by geoai. See :func:plan_geoai_windows.
  • overlap: window overlap in pixels (default: half the window).
  • batch_size: windows per forward pass (default 4; 1 when a raster side is shorter than the window, see :func:plan_geoai_windows).
  • confidence_threshold (default 0.5) and nms_threshold (default 0.3, geoai's cross-window NMS).
  • merge_window_seams (default True), seam_min_px (default 16) and seam_max_gap_px (default 2): join instances that a field was split into at the window edges and fill thin gaps along those edges (see :func:merge_window_seams).
  • clean_instance_mask: False (default), True or a dict of keyword arguments for geoai.utils.raster.clean_instance_mask (removes small instances, fills holes, smooths boundaries; runtime grows with instances x pixels).
required

Returns:

Type Description
GeoDataFrame

One polygon per detected field with instance_id and score columns and gdf.attrs["engine_meta"].

Notes

geoai does not expose two limits of torchvision's Mask R-CNN (:func:maskrcnn_limits, recorded in engine_meta): at most 100 detections are kept per window (box_detections_per_img), so where more fields fit in one window (small fields at 10-30 m) the rest are lost, and detections scoring below 0.05 (box_score_thresh) are discarded, so a confidence_threshold below 0.05 has the same effect as 0.05 (logged at WARNING). For dense small fields, fine-tune with a smaller engine_params["chip_size"]; inference then uses windows of the same size.

Source code in agribound/engines/geoai_field.py
def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
    """Run geoai instance segmentation on a composite.

    Parameters
    ----------
    raster_path : str
        Composite GeoTIFF.
    config : AgriboundConfig
        Pipeline configuration. ``engine_params``:

        - ``checkpoint_path``: Mask R-CNN weights (``.pth``); see
          :func:`resolve_geoai_checkpoint` for the Hugging Face option.
        - ``window_size``: sliding window in pixels (default: the
          training chip size recorded next to the checkpoint, else
          geoai's 512). Keep it equal to the training chip size: Mask
          R-CNN resizes each window to an 800 px shorter side, so the
          window size sets the apparent field size (a different value is
          logged at WARNING). Windows that extend past a small raster are
          zero-padded by geoai. See :func:`plan_geoai_windows`.
        - ``overlap``: window overlap in pixels (default: half the
          window).
        - ``batch_size``: windows per forward pass (default 4; 1 when a
          raster side is shorter than the window, see
          :func:`plan_geoai_windows`).
        - ``confidence_threshold`` (default 0.5) and ``nms_threshold``
          (default 0.3, geoai's cross-window NMS).
        - ``merge_window_seams`` (default *True*), ``seam_min_px``
          (default 16) and ``seam_max_gap_px`` (default 2): join
          instances that a field was split into at the window edges and
          fill thin gaps along those edges (see
          :func:`merge_window_seams`).
        - ``clean_instance_mask``: *False* (default), *True* or a dict of
          keyword arguments for ``geoai.utils.raster.clean_instance_mask``
          (removes small instances, fills holes, smooths boundaries;
          runtime grows with instances x pixels).

    Returns
    -------
    geopandas.GeoDataFrame
        One polygon per detected field with ``instance_id`` and ``score``
        columns and ``gdf.attrs["engine_meta"]``.

    Notes
    -----
    geoai does not expose two limits of torchvision's Mask R-CNN
    (:func:`maskrcnn_limits`, recorded in ``engine_meta``): at most 100
    detections are kept per window (``box_detections_per_img``), so
    where more fields fit in one window (small fields at 10-30 m) the
    rest are lost, and detections scoring below 0.05
    (``box_score_thresh``) are discarded, so a ``confidence_threshold``
    below 0.05 has the same effect as 0.05 (logged at WARNING). For
    dense small fields, fine-tune with a smaller
    ``engine_params["chip_size"]``; inference then uses windows of the
    same size.
    """
    try:
        import torch  # noqa: F401
        from geoai.train import instance_segmentation
    except ImportError:
        raise ImportError(
            "geoai-py is required for the GeoAI engine. "
            "Install with: pip install agribound[geoai]"
        ) from None
    import rasterio

    from agribound._cache import cache_path
    from agribound.engines.base import get_canonical_band_indices
    from agribound.engines.finetune._data import (
        file_sha256,
        mask_invalid_predictions,
        package_versions,
        raster_fingerprint,
        read_json,
        read_training_meta,
        write_json,
        write_rgb_input,
    )

    self.validate_input(raster_path, config)
    params = config.engine_params
    checkpoint, hub_info = resolve_geoai_checkpoint(params)
    indices = get_canonical_band_indices(config.source, ["R", "G", "B"], bands=config.bands)
    confidence = float(params.get("confidence_threshold", 0.5))
    nms = float(params.get("nms_threshold", 0.3))
    limits = maskrcnn_limits()
    if confidence < limits["box_score_thresh"]:
        logger.warning(
            "GeoAI confidence_threshold %.3g is below Mask R-CNN's box_score_thresh %.3g, "
            "which discards lower-scoring detections first; it acts as %.3g",
            confidence,
            limits["box_score_thresh"],
            limits["box_score_thresh"],
        )
    training = read_training_meta(checkpoint)
    with rasterio.open(raster_path) as src:
        height, width = src.height, src.width
    plan = plan_geoai_windows(
        height,
        width,
        window_size=params.get("window_size"),
        overlap=params.get("overlap"),
        batch_size=int(params.get("batch_size", 4)),
        training_chip_size=training.get("chip_size"),
        min_size=int(limits["min_size"]),
    )
    window, overlap, batch_size = plan["window_size"], plan["overlap"], plan["batch_size"]
    clean = params.get("clean_instance_mask", False)

    device = config.resolve_device()
    if device == "mps":
        logger.warning(
            "GeoAI Mask R-CNN runs on CPU instead of MPS: torchvision Mask R-CNN on MPS "
            "reports Metal command-buffer errors and its detections differ from CPU "
            "(checked with torch 2.10 and geoai-py 0.43.1)"
        )
        device = "cpu"

    raster_fp = raster_fingerprint(raster_path)
    rgb_path = cache_path(config, "geoai_rgb", ".tif", raster_fp, f"bands={indices}", "uint8")
    rgb_info_path = rgb_path.with_suffix(".json")
    if rgb_path.exists() and rgb_info_path.exists():
        rgb_info = read_json(rgb_info_path)
    else:
        rgb_info = write_rgb_input(raster_path, rgb_path, indices, unit_float=False)
        write_json(rgb_info_path, rgb_info)

    inst_path = cache_path(
        config,
        "geoai_instances",
        ".tif",
        raster_fp,
        raster_fingerprint(checkpoint),
        f"bands={indices}",
        f"window={window}",
        f"overlap={overlap}",
        f"conf={confidence}",
        f"nms={nms}",
    )
    score_path = inst_path.with_name(f"{inst_path.stem}_score{inst_path.suffix}")
    inst_info_path = inst_path.with_suffix(".json")
    run_info: dict[str, Any] = {"device": device}
    if inst_path.exists() and score_path.exists() and inst_info_path.exists():
        logger.info("Using cached GeoAI instances: %s", inst_path)
        run_info = {**read_json(inst_info_path), "cache_reused": True}
    else:
        logger.info("Running GeoAI instance segmentation (device=%s)", device)
        instance_segmentation(
            input_path=str(rgb_path),
            output_path=str(inst_path),
            model_path=checkpoint,
            window_size=window,
            overlap=overlap,
            confidence_threshold=confidence,
            nms_threshold=nms,
            batch_size=batch_size,
            num_channels=3,
            num_classes=2,
            vectorize=False,
            device=device,
        )
        mask_invalid_predictions(inst_path, raster_path, indices)
        write_json(inst_info_path, run_info)

    vector_source = str(inst_path)
    seam_info: dict[str, Any] = {"seam_merge": False}
    if params.get("merge_window_seams", True):
        seam_min = int(params.get("seam_min_px", 16))
        seam_gap = int(params.get("seam_max_gap_px", 2))
        merged_path = inst_path.with_name(
            f"{inst_path.stem}_seams{seam_min}_gap{seam_gap}{inst_path.suffix}"
        )
        seam_json = merged_path.with_suffix(".json")
        if merged_path.exists() and seam_json.exists():
            seam_info = read_json(seam_json)
        else:
            seam_info = merge_window_seams(
                str(inst_path),
                str(merged_path),
                window,
                overlap,
                min_seam_px=seam_min,
                max_gap_px=seam_gap,
            )
            write_json(seam_json, seam_info)
        logger.info(
            "GeoAI: joined %d instances split at the %d px window edges",
            seam_info["n_instances_merged_at_seams"],
            window,
        )
        vector_source = str(merged_path)
    if clean:
        from geoai.utils.raster import clean_instance_mask

        kwargs = dict(clean) if isinstance(clean, dict) else {}
        # Clean the seam-merged raster when there is one, so the join is kept.
        src = Path(vector_source)
        cleaned = src.with_name(f"{src.stem}_cleaned{src.suffix}")
        vector_source = clean_instance_mask(str(src), str(cleaned), **kwargs)

    gdf = instances_to_polygons(vector_source, str(score_path))
    gdf.attrs["engine_meta"] = {
        "backend": "geoai.instance_segmentation",
        **package_versions("geoai-py", "torch", "torchvision"),
        "model": "maskrcnn_resnet50_fpn",
        "num_classes": 2,
        "checkpoint": str(Path(checkpoint).resolve()),
        "checkpoint_sha256": file_sha256(checkpoint),
        **({"hub": hub_info} if hub_info else {}),
        "training": training or None,
        "band_indices": indices,
        "input": rgb_info,
        "window_size": window,
        "overlap": overlap,
        "window_source": plan["window_source"],
        "training_chip_size": plan["training_chip_size"],
        "resize_factor": plan["resize_factor"],
        "confidence_threshold": confidence,
        "nms_threshold": nms,
        "batch_size": batch_size,
        "requested_batch_size": plan["requested_batch_size"],
        "maskrcnn_limits": limits,
        "clean_instance_mask": clean,
        **seam_info,
        **run_info,
    }
    if len(gdf) == 0:
        logger.warning("No field boundaries detected by GeoAI")
    logger.info("GeoAI delineated %d field boundaries", len(gdf))
    return gdf
prefetch classmethod
prefetch(config: AgriboundConfig) -> list[str]

Download the weights GeoAI inference and fine-tuning need offline.

  • torchvision's COCO Mask R-CNN ResNet50-FPN weights (MaskRCNN_ResNet50_FPN_Weights.DEFAULT): geoai builds its model with them before loading a checkpoint, and fine-tuning starts from them. They go to $TORCH_HOME/hub/checkpoints.
  • The Hugging Face checkpoint, when engine_params["repo_id"] is set.

Returns:

Type Description
list[str]

Local paths.

Source code in agribound/engines/geoai_field.py
@classmethod
def prefetch(cls, config: AgriboundConfig) -> list[str]:
    """Download the weights GeoAI inference and fine-tuning need offline.

    - torchvision's COCO Mask R-CNN ResNet50-FPN weights
      (``MaskRCNN_ResNet50_FPN_Weights.DEFAULT``): geoai builds its model
      with them before loading a checkpoint, and fine-tuning starts from
      them. They go to ``$TORCH_HOME/hub/checkpoints``.
    - The Hugging Face checkpoint, when ``engine_params["repo_id"]`` is
      set.

    Returns
    -------
    list[str]
        Local paths.
    """
    import os

    import torch
    from torchvision.models.detection import MaskRCNN_ResNet50_FPN_Weights

    weights = MaskRCNN_ResNet50_FPN_Weights.DEFAULT
    weights.get_state_dict(progress=True)
    paths = [os.path.join(torch.hub.get_dir(), "checkpoints", os.path.basename(weights.url))]
    params = config.engine_params
    if params.get("repo_id"):
        paths.append(resolve_geoai_checkpoint(params)[0])
    elif params.get("checkpoint_path") and Path(params["checkpoint_path"]).is_file():
        paths.append(str(Path(params["checkpoint_path"]).resolve()))
    return paths

merge_window_seams

merge_window_seams(instance_path: str, output_path: str, window: int, overlap: int, min_seam_px: int = 16, min_seam_fraction: float = 0.5, max_gap_px: int = 2) -> dict[str, Any]

Join instances that one field split into at geoai's window edges.

geoai paints every detection's full mask into one instance raster and keeps the partial detections of a field from overlapping windows when their boxes overlap by less than the NMS threshold, so a field larger than the overlap is split along a window edge (an axis-aligned line at a window start or end), sometimes with a thin gap of background where neither partial mask reaches the edge. For every interior window edge the nearest instances on its two sides are compared row by row (at most max_gap_px background pixels between them): two different instances that meet across the edge along at least min_seam_px pixels, and along at least min_seam_fraction of the shorter of their two runs on that edge, are joined (union-find, so a field cut by several edges becomes one instance). Gaps of at most max_gap_px pixels across an edge between two parts of one (joined) instance are then filled, so no slit is left. Instances that meet anywhere else are left alone.

Writes the relabelled raster to output_path and returns counts for engine_meta.

Source code in agribound/engines/geoai_field.py
def merge_window_seams(
    instance_path: str,
    output_path: str,
    window: int,
    overlap: int,
    min_seam_px: int = 16,
    min_seam_fraction: float = 0.5,
    max_gap_px: int = 2,
) -> dict[str, Any]:
    """Join instances that one field split into at geoai's window edges.

    geoai paints every detection's full mask into one instance raster and keeps the
    partial detections of a field from overlapping windows when their boxes overlap by less
    than the NMS threshold, so a field larger than the overlap is split along a window edge
    (an axis-aligned line at a window start or end), sometimes with a thin gap of
    background where neither partial mask reaches the edge. For every interior window edge
    the nearest instances on its two sides are compared row by row (at most *max_gap_px*
    background pixels between them): two different instances that meet across the edge
    along at least *min_seam_px* pixels, and along at least *min_seam_fraction* of the
    shorter of their two runs on that edge, are joined (union-find, so a field cut by
    several edges becomes one instance). Gaps of at most *max_gap_px* pixels across an
    edge between two parts of one (joined) instance are then filled, so no slit is left.
    Instances that meet anywhere else are left alone.

    Writes the relabelled raster to *output_path* and returns counts for ``engine_meta``.
    """
    import rasterio
    from rasterio.windows import Window

    g = max(int(max_gap_px), 0)
    parent: dict[int, int] = {}

    def find(a: int) -> int:
        while parent.get(a, a) != a:
            parent[a] = parent.get(parent[a], parent[a])
            a = parent[a]
        return a

    def band_windows(width: int, height: int):
        """(window, transpose) for every interior window edge, both axes."""
        for x in window_edges(width, window, overlap):
            lo, hi = max(x - 1 - g, 0), min(x + 1 + g, width)
            yield Window(lo, 0, hi - lo, height), False, x - lo - 1
        for y in window_edges(height, window, overlap):
            lo, hi = max(y - 1 - g, 0), min(y + 1 + g, height)
            yield Window(0, lo, width, hi - lo), True, y - lo - 1

    def as_band(arr: np.ndarray, transpose: bool, before: int) -> tuple[np.ndarray, int]:
        band = arr.T if transpose else arr
        # Re-centre so that the seam lies between columns gg and gg + 1.
        gg = min(before, band.shape[1] - before - 2)
        return band[:, before - gg : before + gg + 2], gg

    with rasterio.open(instance_path) as src:
        profile = src.profile.copy()
        width, height = src.width, src.height
        for win, transpose, before in band_windows(width, height):
            band, gg = as_band(src.read(1, window=win), transpose, before)
            id_b, db, id_a, da = _nearest_across(band, gg)
            meet = (id_b > 0) & (id_a > 0) & (id_b != id_a) & (db + da <= g)
            if not meet.any():
                continue
            key = np.stack(
                [np.minimum(id_b[meet], id_a[meet]), np.maximum(id_b[meet], id_a[meet])], 1
            )
            uniq, counts = np.unique(key, axis=0, return_counts=True)
            run = np.bincount(np.r_[id_b[id_b > 0], id_a[id_a > 0]])
            for (i, j), n in zip(uniq.tolist(), counts.tolist(), strict=True):
                shorter = min(run[i], run[j])
                if n >= min_seam_px and n >= min_seam_fraction * shorter:
                    ri, rj = find(i), find(j)
                    if ri != rj:
                        parent[max(ri, rj)] = min(ri, rj)

        merged_ids = {k for k in parent if find(k) != k}
        lut = None
        if parent:
            lut = np.arange(max(max(parent), max(parent.values())) + 1, dtype=np.int64)
            for k in list(parent):
                lut[k] = find(k)
        with rasterio.open(output_path, "w", **profile) as dst:
            for _, win in src.block_windows(1):
                block = src.read(1, window=win)
                if lut is not None:
                    inside = (block > 0) & (block < len(lut))
                    block = block.copy()
                    block[inside] = lut[block[inside]].astype(block.dtype)
                dst.write(block, 1, window=win)

    n_filled = 0
    if g > 0:
        with rasterio.open(output_path, "r+") as dst:
            for win, transpose, before in band_windows(width, height):
                arr = dst.read(1, window=win)
                band, gg = as_band(arr, transpose, before)
                id_b, db, id_a, da = _nearest_across(band, gg)
                fill = (id_b > 0) & (id_b == id_a) & (db + da > 0) & (db + da <= g)
                if not fill.any():
                    continue
                band = band.copy()
                for r in np.flatnonzero(fill):
                    band[r, gg - db[r] + 1 : gg + 1 + da[r]] = id_b[r]
                    n_filled += int(db[r] + da[r])
                full = arr.T.copy() if transpose else arr.copy()
                full[:, before - gg : before + gg + 2] = band
                dst.write(full.T if transpose else full, 1, window=win)
    return {
        "seam_merge": True,
        "seam_min_px": int(min_seam_px),
        "seam_min_fraction": float(min_seam_fraction),
        "seam_max_gap_px": g,
        "n_instances_merged_at_seams": len(merged_ids),
        "n_seam_gap_pixels_filled": n_filled,
    }

window_edges

window_edges(size: int, window: int, overlap: int) -> list[int]

Pixel offsets of the interior window edges geoai's sliding window uses on one axis.

geoai 0.43.1 places windows at min(k * (window - overlap), size - window) (at least 0) for k = 0 .. ceil((size - overlap) / (window - overlap)); each window spans [start, start + window). The raster borders (0 and size) are left out.

Source code in agribound/engines/geoai_field.py
def window_edges(size: int, window: int, overlap: int) -> list[int]:
    """Pixel offsets of the interior window edges geoai's sliding window uses on one axis.

    geoai 0.43.1 places windows at ``min(k * (window - overlap), size - window)`` (at least
    0) for ``k = 0 .. ceil((size - overlap) / (window - overlap))``; each window spans
    ``[start, start + window)``. The raster borders (0 and *size*) are left out.
    """
    stride = window - overlap
    steps = math.ceil((size - overlap) / stride)
    starts = {max(0, min(k * stride, size - window)) for k in range(steps + 1)}
    edges = starts | {s + window for s in starts}
    return sorted(e for e in edges if 0 < e < size)

DINOv3

dinov3

DINOv3 semantic segmentation engine (geoai-py).

Runs geoai's DINOv3Segmenter -- a DINOv3 ViT backbone (Siméoni et al., 2025, arXiv:2508.10104) with a DPT decoder -- trained by agribound on reference boundaries (fine_tune=True, see agribound.engines.finetune._dinov3) into background / field interior / field boundary classes. There are no published field-boundary weights, so a fine-tuned Lightning .ckpt is required. Each field interior region is grown back over the predicted boundary class by the boundary width used in training (:func:agribound.engines.finetune._data.interior_polygons), so neighbouring fields do not overlap.

Input

Canonical R, G, B bands with a scene-level 1-99 percentile stretch to uint8 (:func:agribound.engines.finetune._data.write_rgb_input), stored as float32 uint8 / 255. For source="local" without config.bands bands 1, 2, 3 are read as R, G, B. geoai divides a window by 255 only when its maximum exceeds 1 and applies no mean/std normalisation, so the model sees exactly these [0, 1] values at training and inference. The SAT-493M backbone was pre-trained with the normalisation mean (0.430, 0.411, 0.296) and standard deviation (0.213, 0.156, 0.143) (facebookresearch/dinov3 README), so its pre-trained features receive inputs that are not normalised as in pre-training; this matters most when the backbone is frozen (use_lora or freeze_backbone). Normalising the chips beforehand is not a workaround: geoai divides any chip or window whose maximum exceeds 1 by 255.

Sliding window

geoai's dinov3_segment_geotiff zero-pads every window that extends past the raster to the full window size, and its last window along each axis starts at min(i * stride, size - 1), so it can extend past the raster. The ViT attends over the whole window, so zero padding changes the features of the real pixels. Agribound therefore (:func:plan_dinov3_windows) uses the training chip size as the window by default, caps the window at the larger raster side (rounded up to the 16 px patch size), and extends the RGB input at the bottom and right by mirror reflection so that every window lies inside it; the prediction is then cropped back to the raster grid.

Weights and offline use

geoai 0.43.1 builds the backbone with torch.hub.load from facebookresearch/dinov3 (GitHub, or the local clone named by the DINOV3_LOCATION environment variable) and then loads the SAT-493M ViT-L/16 weights giswqs/geoai / dinov3_vitl16_sat493m.pth from Hugging Face unless weights_path is given. This also happens when a fine-tuned checkpoint is loaded for inference (with the weights_path recorded in the checkpoint); the checkpoint's weights then replace them. SAT-493M weights exist only for ViT-L/16 and ViT-7B/16, so fine-tuning other backbone sizes needs engine_params["weights_path"]. :meth:DINOv3Engine.prefetch downloads the hub repository and the weights for nodes without internet access.

DINOv3Engine

Bases: DelineationEngine

Field delineation with a fine-tuned DINOv3 + DPT model (see module docstring).

Source code in agribound/engines/dinov3.py
class DINOv3Engine(DelineationEngine):
    """Field delineation with a fine-tuned DINOv3 + DPT model (see module docstring)."""

    name = "dinov3"
    supported_sources = list(ENGINE_REGISTRY["dinov3"]["supported_sources"])
    requires_bands = list(ENGINE_REGISTRY["dinov3"]["requires_bands"])

    def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
        """Run DINOv3 segmentation on a composite.

        Parameters
        ----------
        raster_path : str
            Composite GeoTIFF.
        config : AgriboundConfig
            Pipeline configuration. ``engine_params``:

            - ``checkpoint_path``: Lightning ``.ckpt`` written by
              ``train_dinov3_segmentation`` (required; set automatically by
              ``fine_tune=True``).
            - ``dinov3_model``: only used to check the checkpoint; the
              backbone recorded in the checkpoint is what runs.
            - ``weights_path``: not used for inference. geoai rebuilds the
              backbone from the ``weights_path`` recorded in the checkpoint
              and then loads the fine-tuned weights; a different value here
              only logs a warning.
            - ``window_size``: sliding window in pixels, a multiple of 16
              (default: the training chip size recorded next to the
              checkpoint, else 512; capped at the larger raster side, see
              :func:`plan_dinov3_windows`).
            - ``overlap``: window overlap in pixels (default: half the
              window).
            - ``batch_size``: windows per forward pass (default 4).
            - ``dilate_interior_px``: pixels by which each interior region is
              grown over the predicted boundary class (default: the training
              ``boundary_erosion`` recorded next to the checkpoint, else
              ``engine_params["boundary_erosion"]`` or 2).

        Returns
        -------
        geopandas.GeoDataFrame
            Field polygons with ``gdf.attrs["engine_meta"]``.

        Raises
        ------
        RuntimeError
            Without a checkpoint, or if geoai writes no output.
        ValueError
            If the checkpoint is not a Lightning ``.ckpt``.
        """
        try:
            from geoai.dinov3_finetune import dinov3_segment_geotiff
        except ImportError:
            raise ImportError(
                "geoai-py is required for the DINOv3 engine. "
                "Install with: pip install agribound[dinov3]"
            ) from None

        import rasterio

        from agribound._cache import cache_path
        from agribound.engines.base import get_canonical_band_indices
        from agribound.engines.finetune._data import (
            file_sha256,
            interior_polygons,
            mask_invalid_predictions,
            package_versions,
            raster_fingerprint,
            read_checkpoint_hparams,
            read_json,
            read_training_meta,
            write_json,
            write_rgb_input,
        )

        self.validate_input(raster_path, config)
        params = config.engine_params
        checkpoint = params.get("checkpoint_path")
        if not checkpoint:
            raise RuntimeError(
                "DINOv3 requires a fine-tuned checkpoint (no field-boundary weights are "
                "published). Set fine_tune=True with reference_boundaries, or provide "
                "engine_params={'checkpoint_path': '/path/to/dinov3.ckpt'}."
            )
        ckpt = Path(checkpoint).expanduser()
        if not ckpt.is_file():
            raise FileNotFoundError(f"DINOv3 checkpoint not found: {ckpt}")
        if ckpt.suffix != ".ckpt":
            raise ValueError(
                f"DINOv3 checkpoint {ckpt} is not a Lightning .ckpt. geoai loads other files "
                "as a plain state dict with strict=False, which can silently skip weights; "
                "use the .ckpt written by fine-tuning."
            )

        hparams = read_checkpoint_hparams(ckpt)
        model_name = hparams.get("model_name", DINOV3_DEFAULT_BACKBONE)
        requested = params.get("dinov3_model")
        if requested is not None:
            wanted = DINOV3_MODELS.get(str(requested), str(requested))
            if wanted != model_name:
                logger.warning(
                    "engine_params['dinov3_model']=%r differs from the checkpoint's backbone "
                    "%r; the checkpoint's backbone is used",
                    requested,
                    model_name,
                )
        # geoai rebuilds the model from the checkpoint's hyper-parameters
        # (DINOv3Segmenter.load_from_checkpoint): the backbone is initialised
        # from hparams["weights_path"] (the SAT-493M ViT-L/16 file when it is
        # None or missing), then the checkpoint's state dict replaces every
        # weight (strict loading).
        weights_path = hparams.get("weights_path")
        user_weights = params.get("weights_path")
        if user_weights and str(user_weights) != str(weights_path):
            logger.warning(
                "engine_params['weights_path']=%r is not used for inference: geoai rebuilds "
                "the backbone from the checkpoint's own weights_path (%r) before loading the "
                "fine-tuned weights",
                user_weights,
                weights_path,
            )
        if weights_path and not Path(weights_path).is_file():
            if model_name != DINOV3_DEFAULT_BACKBONE:
                raise FileNotFoundError(
                    f"The {model_name} checkpoint {ckpt} was trained from backbone weights "
                    f"{weights_path!r}, which do not exist on this machine. geoai would "
                    "initialise the backbone from the ViT-L/16 SAT-493M weights instead, which "
                    f"do not fit {model_name}. Copy the weights file to that path."
                )
            logger.warning(
                "Backbone weights %r recorded in %s do not exist here; geoai initialises the "
                "ViT-L/16 backbone from %s/%s instead. The fine-tuned checkpoint then replaces "
                "every weight, so the prediction is unchanged.",
                weights_path,
                ckpt.name,
                *DINOV3_DEFAULT_WEIGHTS,
            )
            weights_path = None
        batch_size = int(params.get("batch_size", 4))
        training = read_training_meta(ckpt)
        dilate_px = params.get("dilate_interior_px")
        if dilate_px is None:
            dilate_px = training.get("boundary_erosion", params.get("boundary_erosion", 2))
        dilate_px = int(dilate_px)
        with rasterio.open(raster_path) as src:
            height, width = src.height, src.width
        plan = plan_dinov3_windows(
            height,
            width,
            window_size=params.get("window_size"),
            overlap=params.get("overlap"),
            training_chip_size=training.get("chip_size"),
        )
        window_size, overlap = plan["window_size"], plan["overlap"]
        padded = (plan["padded_height"], plan["padded_width"])
        if plan["capped"]:
            logger.info(
                "DINOv3 window capped at %d px (raster %d x %d px)", window_size, height, width
            )
        if plan["training_chip_size"] and window_size != plan["training_chip_size"]:
            logger.info(
                "DINOv3 window %d px differs from the %d px training chips (the attention "
                "context differs; the object scale does not)",
                window_size,
                plan["training_chip_size"],
            )
        indices = get_canonical_band_indices(config.source, ["R", "G", "B"], bands=config.bands)
        device = config.resolve_device()
        raster_fp = raster_fingerprint(raster_path)

        rgb_path = cache_path(
            config,
            "dinov3_rgb",
            ".tif",
            raster_fp,
            f"bands={indices}",
            "unit",
            f"padded={padded[0]}x{padded[1]}:reflect",
        )
        rgb_info_path = rgb_path.with_suffix(".json")
        if rgb_path.exists() and rgb_info_path.exists():
            rgb_info = read_json(rgb_info_path)
        else:
            rgb_info = write_rgb_input(
                raster_path, rgb_path, indices, unit_float=True, pad_to=padded
            )
            write_json(rgb_info_path, rgb_info)

        seg_path = cache_path(
            config,
            "dinov3_segmentation",
            ".tif",
            raster_fp,
            raster_fingerprint(ckpt),
            f"bands={indices}",
            f"window={window_size}",
            f"overlap={overlap}",
            f"padded={padded[0]}x{padded[1]}:reflect",
            f"weights={weights_path}",
        )
        seg_info_path = seg_path.with_suffix(".json")
        run_info = {"device": device}
        if seg_path.exists() and seg_info_path.exists():
            logger.info("Using cached DINOv3 segmentation: %s", seg_path)
            run_info = {**read_json(seg_info_path), "cache_reused": True}
        else:
            logger.info(
                "Running DINOv3 segmentation (backbone=%s, checkpoint=%s, window=%d, "
                "overlap=%d, device=%s)",
                model_name,
                ckpt,
                window_size,
                overlap,
                device,
            )
            raw = seg_path.with_name(seg_path.stem + ".padded.partial.tif")
            tmp = seg_path.with_name(seg_path.stem + ".partial.tif")
            dinov3_segment_geotiff(
                input_path=str(rgb_path),
                output_path=str(raw),
                checkpoint_path=str(ckpt),
                model_name=model_name,
                weights_path=weights_path,
                num_classes=int(hparams.get("num_classes", 3)),
                window_size=window_size,
                overlap=overlap,
                batch_size=batch_size,
                device=device,
                quiet=True,
            )
            if not raw.exists():
                raise RuntimeError(f"DINOv3 segmentation failed: no output at {raw}")
            if padded != (height, width):
                _crop_raster(raw, tmp, height, width)
                raw.unlink()
            else:
                raw.replace(tmp)
            mask_invalid_predictions(tmp, raster_path, indices)
            tmp.replace(seg_path)
            write_json(seg_info_path, run_info)

        gdf = interior_polygons(seg_path, dilate_px, config.min_field_area_m2)
        gdf.attrs["engine_meta"] = {
            "backend": "geoai",
            **package_versions("geoai-py", "torch"),
            "model_name": model_name,
            # Pre-trained backbone weights the checkpoint was fine-tuned from.
            "weights": training.get("weights")
            or hparams.get("weights_path")
            or "/".join(DINOV3_DEFAULT_WEIGHTS),
            "weights_revision": training.get("weights_revision"),
            "weights_sha256": training.get("weights_sha256"),
            "use_lora": hparams.get("use_lora"),
            "freeze_backbone": hparams.get("freeze_backbone"),
            "lora_rank": hparams.get("lora_rank") if hparams.get("use_lora") else None,
            "trainable_params": training.get("trainable_params"),
            "checkpoint": str(ckpt.resolve()),
            "checkpoint_sha256": file_sha256(ckpt),
            "band_indices": indices,
            "input": rgb_info,
            "window_size": window_size,
            "overlap": overlap,
            "window_source": plan["window_source"],
            "window_capped": plan["capped"],
            "training_chip_size": plan["training_chip_size"],
            "input_padded_to": list(padded),
            "input_padding": "reflect" if padded != (height, width) else None,
            "batch_size": batch_size,
            "field_class": 1,
            "dilate_interior_px": dilate_px,
            "dinov3_location": os.environ.get("DINOV3_LOCATION", "facebookresearch/dinov3"),
            **run_info,
        }
        logger.info("DINOv3 delineated %d field boundaries", len(gdf))
        return gdf

    @classmethod
    def prefetch(cls, config: AgriboundConfig) -> list[str]:
        """Download what geoai needs to build the DINOv3 backbone offline.

        - The ``facebookresearch/dinov3`` torch.hub repository (skipped when
          ``DINOV3_LOCATION`` points to a local clone). On nodes without
          internet access set ``DINOV3_LOCATION`` to the returned directory.
        - The SAT-493M ViT-L/16 weights from Hugging Face (skipped when
          ``engine_params["weights_path"]`` is given); set
          ``HF_HUB_OFFLINE=1`` on offline nodes.

        Returns
        -------
        list[str]
            Local paths (hub repository directory, weights file).
        """
        import torch
        from huggingface_hub import hf_hub_download

        paths: list[str] = []
        location = os.environ.get("DINOV3_LOCATION")
        if location and Path(location).is_dir():
            paths.append(str(Path(location).resolve()))
        else:
            torch.hub.list("facebookresearch/dinov3", trust_repo=True, skip_validation=True)
            hub_dir = Path(torch.hub.get_dir())
            candidates = sorted(hub_dir.glob("facebookresearch_dinov3_*"))
            if not candidates:
                raise RuntimeError(f"torch.hub did not cache facebookresearch/dinov3 in {hub_dir}")
            main = hub_dir / "facebookresearch_dinov3_main"
            repo_dir = main if main in candidates else candidates[0]
            paths.append(str(repo_dir))
            logger.info("Cached facebookresearch/dinov3 in %s (use as DINOV3_LOCATION)", repo_dir)
        weights_path = config.engine_params.get("weights_path")
        if weights_path:
            paths.append(str(Path(weights_path).resolve()))
        else:
            paths.append(
                hf_hub_download(
                    repo_id=DINOV3_DEFAULT_WEIGHTS[0], filename=DINOV3_DEFAULT_WEIGHTS[1]
                )
            )
        return paths
delineate
delineate(raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame

Run DINOv3 segmentation on a composite.

Parameters:

Name Type Description Default
raster_path str

Composite GeoTIFF.

required
config AgriboundConfig

Pipeline configuration. engine_params:

  • checkpoint_path: Lightning .ckpt written by train_dinov3_segmentation (required; set automatically by fine_tune=True).
  • dinov3_model: only used to check the checkpoint; the backbone recorded in the checkpoint is what runs.
  • weights_path: not used for inference. geoai rebuilds the backbone from the weights_path recorded in the checkpoint and then loads the fine-tuned weights; a different value here only logs a warning.
  • window_size: sliding window in pixels, a multiple of 16 (default: the training chip size recorded next to the checkpoint, else 512; capped at the larger raster side, see :func:plan_dinov3_windows).
  • overlap: window overlap in pixels (default: half the window).
  • batch_size: windows per forward pass (default 4).
  • dilate_interior_px: pixels by which each interior region is grown over the predicted boundary class (default: the training boundary_erosion recorded next to the checkpoint, else engine_params["boundary_erosion"] or 2).
required

Returns:

Type Description
GeoDataFrame

Field polygons with gdf.attrs["engine_meta"].

Raises:

Type Description
RuntimeError

Without a checkpoint, or if geoai writes no output.

ValueError

If the checkpoint is not a Lightning .ckpt.

Source code in agribound/engines/dinov3.py
def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
    """Run DINOv3 segmentation on a composite.

    Parameters
    ----------
    raster_path : str
        Composite GeoTIFF.
    config : AgriboundConfig
        Pipeline configuration. ``engine_params``:

        - ``checkpoint_path``: Lightning ``.ckpt`` written by
          ``train_dinov3_segmentation`` (required; set automatically by
          ``fine_tune=True``).
        - ``dinov3_model``: only used to check the checkpoint; the
          backbone recorded in the checkpoint is what runs.
        - ``weights_path``: not used for inference. geoai rebuilds the
          backbone from the ``weights_path`` recorded in the checkpoint
          and then loads the fine-tuned weights; a different value here
          only logs a warning.
        - ``window_size``: sliding window in pixels, a multiple of 16
          (default: the training chip size recorded next to the
          checkpoint, else 512; capped at the larger raster side, see
          :func:`plan_dinov3_windows`).
        - ``overlap``: window overlap in pixels (default: half the
          window).
        - ``batch_size``: windows per forward pass (default 4).
        - ``dilate_interior_px``: pixels by which each interior region is
          grown over the predicted boundary class (default: the training
          ``boundary_erosion`` recorded next to the checkpoint, else
          ``engine_params["boundary_erosion"]`` or 2).

    Returns
    -------
    geopandas.GeoDataFrame
        Field polygons with ``gdf.attrs["engine_meta"]``.

    Raises
    ------
    RuntimeError
        Without a checkpoint, or if geoai writes no output.
    ValueError
        If the checkpoint is not a Lightning ``.ckpt``.
    """
    try:
        from geoai.dinov3_finetune import dinov3_segment_geotiff
    except ImportError:
        raise ImportError(
            "geoai-py is required for the DINOv3 engine. "
            "Install with: pip install agribound[dinov3]"
        ) from None

    import rasterio

    from agribound._cache import cache_path
    from agribound.engines.base import get_canonical_band_indices
    from agribound.engines.finetune._data import (
        file_sha256,
        interior_polygons,
        mask_invalid_predictions,
        package_versions,
        raster_fingerprint,
        read_checkpoint_hparams,
        read_json,
        read_training_meta,
        write_json,
        write_rgb_input,
    )

    self.validate_input(raster_path, config)
    params = config.engine_params
    checkpoint = params.get("checkpoint_path")
    if not checkpoint:
        raise RuntimeError(
            "DINOv3 requires a fine-tuned checkpoint (no field-boundary weights are "
            "published). Set fine_tune=True with reference_boundaries, or provide "
            "engine_params={'checkpoint_path': '/path/to/dinov3.ckpt'}."
        )
    ckpt = Path(checkpoint).expanduser()
    if not ckpt.is_file():
        raise FileNotFoundError(f"DINOv3 checkpoint not found: {ckpt}")
    if ckpt.suffix != ".ckpt":
        raise ValueError(
            f"DINOv3 checkpoint {ckpt} is not a Lightning .ckpt. geoai loads other files "
            "as a plain state dict with strict=False, which can silently skip weights; "
            "use the .ckpt written by fine-tuning."
        )

    hparams = read_checkpoint_hparams(ckpt)
    model_name = hparams.get("model_name", DINOV3_DEFAULT_BACKBONE)
    requested = params.get("dinov3_model")
    if requested is not None:
        wanted = DINOV3_MODELS.get(str(requested), str(requested))
        if wanted != model_name:
            logger.warning(
                "engine_params['dinov3_model']=%r differs from the checkpoint's backbone "
                "%r; the checkpoint's backbone is used",
                requested,
                model_name,
            )
    # geoai rebuilds the model from the checkpoint's hyper-parameters
    # (DINOv3Segmenter.load_from_checkpoint): the backbone is initialised
    # from hparams["weights_path"] (the SAT-493M ViT-L/16 file when it is
    # None or missing), then the checkpoint's state dict replaces every
    # weight (strict loading).
    weights_path = hparams.get("weights_path")
    user_weights = params.get("weights_path")
    if user_weights and str(user_weights) != str(weights_path):
        logger.warning(
            "engine_params['weights_path']=%r is not used for inference: geoai rebuilds "
            "the backbone from the checkpoint's own weights_path (%r) before loading the "
            "fine-tuned weights",
            user_weights,
            weights_path,
        )
    if weights_path and not Path(weights_path).is_file():
        if model_name != DINOV3_DEFAULT_BACKBONE:
            raise FileNotFoundError(
                f"The {model_name} checkpoint {ckpt} was trained from backbone weights "
                f"{weights_path!r}, which do not exist on this machine. geoai would "
                "initialise the backbone from the ViT-L/16 SAT-493M weights instead, which "
                f"do not fit {model_name}. Copy the weights file to that path."
            )
        logger.warning(
            "Backbone weights %r recorded in %s do not exist here; geoai initialises the "
            "ViT-L/16 backbone from %s/%s instead. The fine-tuned checkpoint then replaces "
            "every weight, so the prediction is unchanged.",
            weights_path,
            ckpt.name,
            *DINOV3_DEFAULT_WEIGHTS,
        )
        weights_path = None
    batch_size = int(params.get("batch_size", 4))
    training = read_training_meta(ckpt)
    dilate_px = params.get("dilate_interior_px")
    if dilate_px is None:
        dilate_px = training.get("boundary_erosion", params.get("boundary_erosion", 2))
    dilate_px = int(dilate_px)
    with rasterio.open(raster_path) as src:
        height, width = src.height, src.width
    plan = plan_dinov3_windows(
        height,
        width,
        window_size=params.get("window_size"),
        overlap=params.get("overlap"),
        training_chip_size=training.get("chip_size"),
    )
    window_size, overlap = plan["window_size"], plan["overlap"]
    padded = (plan["padded_height"], plan["padded_width"])
    if plan["capped"]:
        logger.info(
            "DINOv3 window capped at %d px (raster %d x %d px)", window_size, height, width
        )
    if plan["training_chip_size"] and window_size != plan["training_chip_size"]:
        logger.info(
            "DINOv3 window %d px differs from the %d px training chips (the attention "
            "context differs; the object scale does not)",
            window_size,
            plan["training_chip_size"],
        )
    indices = get_canonical_band_indices(config.source, ["R", "G", "B"], bands=config.bands)
    device = config.resolve_device()
    raster_fp = raster_fingerprint(raster_path)

    rgb_path = cache_path(
        config,
        "dinov3_rgb",
        ".tif",
        raster_fp,
        f"bands={indices}",
        "unit",
        f"padded={padded[0]}x{padded[1]}:reflect",
    )
    rgb_info_path = rgb_path.with_suffix(".json")
    if rgb_path.exists() and rgb_info_path.exists():
        rgb_info = read_json(rgb_info_path)
    else:
        rgb_info = write_rgb_input(
            raster_path, rgb_path, indices, unit_float=True, pad_to=padded
        )
        write_json(rgb_info_path, rgb_info)

    seg_path = cache_path(
        config,
        "dinov3_segmentation",
        ".tif",
        raster_fp,
        raster_fingerprint(ckpt),
        f"bands={indices}",
        f"window={window_size}",
        f"overlap={overlap}",
        f"padded={padded[0]}x{padded[1]}:reflect",
        f"weights={weights_path}",
    )
    seg_info_path = seg_path.with_suffix(".json")
    run_info = {"device": device}
    if seg_path.exists() and seg_info_path.exists():
        logger.info("Using cached DINOv3 segmentation: %s", seg_path)
        run_info = {**read_json(seg_info_path), "cache_reused": True}
    else:
        logger.info(
            "Running DINOv3 segmentation (backbone=%s, checkpoint=%s, window=%d, "
            "overlap=%d, device=%s)",
            model_name,
            ckpt,
            window_size,
            overlap,
            device,
        )
        raw = seg_path.with_name(seg_path.stem + ".padded.partial.tif")
        tmp = seg_path.with_name(seg_path.stem + ".partial.tif")
        dinov3_segment_geotiff(
            input_path=str(rgb_path),
            output_path=str(raw),
            checkpoint_path=str(ckpt),
            model_name=model_name,
            weights_path=weights_path,
            num_classes=int(hparams.get("num_classes", 3)),
            window_size=window_size,
            overlap=overlap,
            batch_size=batch_size,
            device=device,
            quiet=True,
        )
        if not raw.exists():
            raise RuntimeError(f"DINOv3 segmentation failed: no output at {raw}")
        if padded != (height, width):
            _crop_raster(raw, tmp, height, width)
            raw.unlink()
        else:
            raw.replace(tmp)
        mask_invalid_predictions(tmp, raster_path, indices)
        tmp.replace(seg_path)
        write_json(seg_info_path, run_info)

    gdf = interior_polygons(seg_path, dilate_px, config.min_field_area_m2)
    gdf.attrs["engine_meta"] = {
        "backend": "geoai",
        **package_versions("geoai-py", "torch"),
        "model_name": model_name,
        # Pre-trained backbone weights the checkpoint was fine-tuned from.
        "weights": training.get("weights")
        or hparams.get("weights_path")
        or "/".join(DINOV3_DEFAULT_WEIGHTS),
        "weights_revision": training.get("weights_revision"),
        "weights_sha256": training.get("weights_sha256"),
        "use_lora": hparams.get("use_lora"),
        "freeze_backbone": hparams.get("freeze_backbone"),
        "lora_rank": hparams.get("lora_rank") if hparams.get("use_lora") else None,
        "trainable_params": training.get("trainable_params"),
        "checkpoint": str(ckpt.resolve()),
        "checkpoint_sha256": file_sha256(ckpt),
        "band_indices": indices,
        "input": rgb_info,
        "window_size": window_size,
        "overlap": overlap,
        "window_source": plan["window_source"],
        "window_capped": plan["capped"],
        "training_chip_size": plan["training_chip_size"],
        "input_padded_to": list(padded),
        "input_padding": "reflect" if padded != (height, width) else None,
        "batch_size": batch_size,
        "field_class": 1,
        "dilate_interior_px": dilate_px,
        "dinov3_location": os.environ.get("DINOV3_LOCATION", "facebookresearch/dinov3"),
        **run_info,
    }
    logger.info("DINOv3 delineated %d field boundaries", len(gdf))
    return gdf
prefetch classmethod
prefetch(config: AgriboundConfig) -> list[str]

Download what geoai needs to build the DINOv3 backbone offline.

  • The facebookresearch/dinov3 torch.hub repository (skipped when DINOV3_LOCATION points to a local clone). On nodes without internet access set DINOV3_LOCATION to the returned directory.
  • The SAT-493M ViT-L/16 weights from Hugging Face (skipped when engine_params["weights_path"] is given); set HF_HUB_OFFLINE=1 on offline nodes.

Returns:

Type Description
list[str]

Local paths (hub repository directory, weights file).

Source code in agribound/engines/dinov3.py
@classmethod
def prefetch(cls, config: AgriboundConfig) -> list[str]:
    """Download what geoai needs to build the DINOv3 backbone offline.

    - The ``facebookresearch/dinov3`` torch.hub repository (skipped when
      ``DINOV3_LOCATION`` points to a local clone). On nodes without
      internet access set ``DINOV3_LOCATION`` to the returned directory.
    - The SAT-493M ViT-L/16 weights from Hugging Face (skipped when
      ``engine_params["weights_path"]`` is given); set
      ``HF_HUB_OFFLINE=1`` on offline nodes.

    Returns
    -------
    list[str]
        Local paths (hub repository directory, weights file).
    """
    import torch
    from huggingface_hub import hf_hub_download

    paths: list[str] = []
    location = os.environ.get("DINOV3_LOCATION")
    if location and Path(location).is_dir():
        paths.append(str(Path(location).resolve()))
    else:
        torch.hub.list("facebookresearch/dinov3", trust_repo=True, skip_validation=True)
        hub_dir = Path(torch.hub.get_dir())
        candidates = sorted(hub_dir.glob("facebookresearch_dinov3_*"))
        if not candidates:
            raise RuntimeError(f"torch.hub did not cache facebookresearch/dinov3 in {hub_dir}")
        main = hub_dir / "facebookresearch_dinov3_main"
        repo_dir = main if main in candidates else candidates[0]
        paths.append(str(repo_dir))
        logger.info("Cached facebookresearch/dinov3 in %s (use as DINOV3_LOCATION)", repo_dir)
    weights_path = config.engine_params.get("weights_path")
    if weights_path:
        paths.append(str(Path(weights_path).resolve()))
    else:
        paths.append(
            hf_hub_download(
                repo_id=DINOV3_DEFAULT_WEIGHTS[0], filename=DINOV3_DEFAULT_WEIGHTS[1]
            )
        )
    return paths

Prithvi-EO-2.0

prithvi

Prithvi-EO-2.0 engine (terratorch).

Prithvi-EO-2.0 (Szwarcman et al., 2026, IEEE TGRS, doi:10.1109/TGRS.2025.3642610) is a ViT masked-autoencoder pre-trained on HLS Blue, Green, Red, narrow NIR, SWIR 1 and SWIR 2 surface reflectance. Agribound builds the encoder from the terratorch backbone registry and runs it on single-date composites (num_frames=1). Three modes (engine_params["mode"]):

"embed" (label-free) Patch-token features of one encoder layer (the last, normalised layer by default) from non-overlapping tiles (tiles that extend past the raster are filled by mirror reflection of the raster), interpolated bilinearly between patch-token centres to pixel resolution and clustered with K-means (fitted on a seeded sample of at most 50 000 pixels); 4-connected regions of one cluster become polygons. Clusters are land-cover segments, not field instances. "segment" A Prithvi + UPerNet segmentation model fine-tuned by agribound (fine_tune=True, see agribound.engines.finetune._prithvi) or any terratorch SemanticSegmentationTask checkpoint trained on the same six bands and normalisation with class 1 = field interior and class 2 = field boundary, run with terratorch's tiled_inference on the whole raster (held in memory with its class logits). Each field interior region (class 1) is grown back over the predicted boundary class (2) by the boundary width used in training, so neighbouring fields do not overlap. "pca" Baseline without the ViT: K-means on the PCA of per-band z-scores of R, G, B, NIR.

The default mode is "segment" when engine_params["checkpoint_path"] is set (as after fine-tuning) and "embed" otherwise.

Inputs are the six bands of :func:agribound.engines.finetune._data.prithvi_band_names in surface reflectance x 10000 (the scale of agribound's Sentinel-2, Landsat and HLS composites), normalised with the Prithvi-EO-2.0 means and standard deviations (:data:PRITHVI_MEAN, :data:PRITHVI_STD); no other scaling is applied. Invalid pixels are set to the band means (0 after normalisation) and to label 0 in the output. The embed and segment modes need terratorch (pip install agribound[prithvi] or environment-gfm.yml).

PrithviEngine

Bases: DelineationEngine

Field delineation with Prithvi-EO-2.0 (see the module docstring).

Source code in agribound/engines/prithvi.py
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
class PrithviEngine(DelineationEngine):
    """Field delineation with Prithvi-EO-2.0 (see the module docstring)."""

    name = "prithvi"
    supported_sources = list(ENGINE_REGISTRY["prithvi"]["supported_sources"])
    requires_bands = list(ENGINE_REGISTRY["prithvi"]["requires_bands"])

    def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
        """Run Prithvi-based field delineation.

        Parameters
        ----------
        raster_path : str
            Composite GeoTIFF.
        config : AgriboundConfig
            Pipeline configuration. ``engine_params``:

            - ``mode``: ``"embed"`` | ``"segment"`` | ``"pca"`` (default: see
              the module docstring).
            - ``checkpoint_path``: terratorch ``SemanticSegmentationTask``
              checkpoint (``segment`` mode).
            - ``model_name``: Prithvi-EO-2.0 variant for ``embed`` mode
              (default ``"Prithvi-EO-2.0-300M-TL"``; see :data:`PRITHVI_MODELS`).
            - ``tile_size``: tile edge in pixels (default 224). In ``embed``
              mode it must be a multiple of the patch size (16; 14 for the
              600M models); tiles that extend past the raster are filled by
              mirror reflection of the raster. In ``segment`` mode a raster
              that fits in one tile is passed whole (terratorch reflect-pads
              it to a multiple of twice the patch size); a larger raster is
              run through terratorch's ``tiled_inference``, after mirror
              padding a side shorter than the tile to the tile size. On
              Apple MPS the segmentation model runs on CPU (logged at
              WARNING) unless the model input size passes
              :func:`upernet_mps_compatible` (e.g. 192 px tiles for patch
              16). ``patch_size`` is accepted as a legacy alias.
            - ``stride``: tile step for ``segment`` mode (default
              ``tile_size - 32``).
            - ``batch_size``: tiles per forward pass (default 8).
            - ``n_clusters``: int or ``"auto"`` (silhouette over 5, 10, 15,
              20, 30 on a seeded sample) for ``embed``/``pca``.
            - ``embed_layer``: encoder layer used in ``embed`` mode (default
              -1, the normalised last layer).
            - ``temporal_coords``: ``[year, day_of_year]``, or *False* to not
              pass them (default: :func:`composite_mid_date`); ``embed``
              mode with a ``*-TL`` model only (``segment`` mode and
              fine-tuning pass no coordinates).
            - ``location_coords``: ``[lat, lon]``, or *False* (default: the
              raster centre); ``embed`` mode with a ``*-TL`` model only.
            - ``dilate_interior_px``: ``segment`` mode; pixels by which each
              interior region is grown over the predicted boundary class
              (default: the checkpoint's training ``boundary_erosion`` if
              recorded, else ``engine_params["boundary_erosion"]`` or 2; see
              :func:`agribound.engines.finetune._data.interior_polygons`).
            - ``value_scale``: required for ``source="local"``
              (``"reflectance_x10000"`` or ``"unit"``). Without
              ``config.bands`` a local raster's bands 1-6 are read as Blue,
              Green, Red, narrow NIR, SWIR 1, SWIR 2 (``pca`` mode: bands
              1-4 as R, G, B, NIR).
            - ``pretrained``: *False* builds a randomly initialised encoder in
              ``embed`` mode (for tests only; logged as a warning).

        Returns
        -------
        geopandas.GeoDataFrame
            Polygons with ``gdf.attrs["engine_meta"]``.

        Raises
        ------
        RuntimeError
            If ``mode="segment"`` is requested without a checkpoint.
        ValueError
            For unknown modes, models or non-reflectance inputs.
        """
        params = config.engine_params
        checkpoint = params.get("checkpoint_path")
        mode = str(params.get("mode") or ("segment" if checkpoint else "embed")).lower()
        if mode not in _MODES:
            raise ValueError(f"Unknown Prithvi mode {mode!r}. Choose from {_MODES}")
        self.validate_input(raster_path, config)
        if mode == "segment":
            if not checkpoint:
                raise RuntimeError(
                    "Prithvi mode='segment' needs a fine-tuned checkpoint: set fine_tune=True "
                    "with reference_boundaries, or engine_params['checkpoint_path']. Use "
                    "mode='embed' for label-free clustering."
                )
            return self._segment_mode(raster_path, config, str(checkpoint))
        if checkpoint:
            logger.warning(
                "Prithvi mode=%r ignores engine_params['checkpoint_path'] (%s)", mode, checkpoint
            )
        if mode == "pca":
            return self._pca_mode(raster_path, config)
        return self._embed_mode(raster_path, config)

    # ------------------------------------------------------------------
    # Input
    # ------------------------------------------------------------------

    @staticmethod
    def _bands(config: AgriboundConfig) -> tuple[list[str], list[int]]:
        from agribound.engines.base import get_canonical_band_indices
        from agribound.engines.finetune._data import prithvi_band_names

        names = prithvi_band_names(config.source, config.bands)
        return names, get_canonical_band_indices(config.source, names, bands=config.bands)

    @staticmethod
    def _normalise(
        data: np.ndarray, source: str, value_scale: str | None, nodata: float | None
    ) -> tuple[np.ndarray, np.ndarray]:
        """Return (normalised float32 bands with invalid = 0, valid mask)."""
        from agribound.engines.finetune._data import prithvi_reflectance, valid_pixels

        valid = valid_pixels(data, nodata)
        refl = prithvi_reflectance(data, source, value_scale)
        mean = np.asarray(PRITHVI_MEAN, dtype=np.float32)[:, None, None]
        std = np.asarray(PRITHVI_STD, dtype=np.float32)[:, None, None]
        norm = (refl - mean) / std
        norm = np.where(valid[None], norm, 0.0).astype(np.float32)
        return norm, valid

    @staticmethod
    def _coords(config: AgriboundConfig, raster_path: str) -> dict[str, Any]:
        params = config.engine_params
        out: dict[str, Any] = {}
        tc = params.get("temporal_coords")
        if tc is not False:
            year, doy = composite_mid_date(config) if tc is None else (tc[0], tc[1])
            out["temporal_coords"] = [float(year), float(doy)]
        lc = params.get("location_coords")
        if lc is not False:
            lat, lon = raster_centre_latlon(raster_path) if lc is None else (lc[0], lc[1])
            out["location_coords"] = [float(lat), float(lon)]
        return out

    # ------------------------------------------------------------------
    # Embed mode
    # ------------------------------------------------------------------

    def _embed_mode(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
        """Label-free clustering of Prithvi patch-token embeddings."""
        import rasterio
        from rasterio.windows import Window

        from agribound._cache import cache_path
        from agribound._repro import get_rng
        from agribound.engines.finetune._data import (
            hf_cached_file_info,
            package_versions,
            raster_fingerprint,
            read_json,
            write_json,
        )

        torch = _import_terratorch("embedding mode")
        from terratorch.registry import BACKBONE_REGISTRY

        params = config.engine_params
        registry_name = resolve_prithvi_model(params.get("model_name"))
        depth, patch = PRITHVI_ARCH[registry_name]
        tile = int(params.get("tile_size", params.get("patch_size", 224)))
        if tile % patch:
            raise ValueError(f"tile_size {tile} must be a multiple of the patch size {patch}")
        batch_size = max(1, int(params.get("batch_size", 8)))
        layer = int(params.get("embed_layer", -1))
        n_clusters = params.get("n_clusters", "auto")
        value_scale = params.get("value_scale")
        pretrained = bool(params.get("pretrained", True))
        is_tl = registry_name.endswith("_tl")
        coords = self._coords(config, raster_path) if is_tl else {}
        names, indices = self._bands(config)
        device = config.resolve_device()

        cluster_path = cache_path(
            config,
            "prithvi_embed_clusters",
            ".tif",
            raster_fingerprint(raster_path),
            registry_name,
            f"pretrained={pretrained}",
            f"bands={names}:{indices}",
            f"value_scale={value_scale}",
            f"tile={tile}",
            "pad=reflect",
            f"layer={layer}",
            f"coords={coords}",
            f"k={n_clusters}",
            f"seed={config.seed}",
        )
        meta: dict[str, Any] = {
            "backend": "terratorch",
            "mode": "embed",
            "model": registry_name,
            "weights_repo": PRITHVI_WEIGHTS[registry_name][0] if pretrained else None,
            "weights_file": PRITHVI_WEIGHTS[registry_name][1] if pretrained else None,
            "pretrained": pretrained,
            "num_frames": 1,
            "band_names": names,
            "band_indices": indices,
            "input_units": "surface reflectance x 10000",
            "normalisation": {"mean": PRITHVI_MEAN, "std": PRITHVI_STD},
            "tile_size": tile,
            "tile_padding": "reflect",
            "patch_size": patch,
            "embed_layer": layer,
            "coords": coords,
            "seed": config.seed,
            "device": device,
        }
        meta.update(package_versions("terratorch", "torch"))
        if config.source == "landsat":
            meta["nir_note"] = "Landsat SR_B5: narrow NIR on L8/9, broad NIR on L5/7"

        sidecar = cluster_path.with_suffix(".json")
        if cluster_path.exists() and sidecar.exists():
            logger.info("Using cached Prithvi clusters: %s", cluster_path)
            meta.update(read_json(sidecar), cache_reused=True)
        else:
            if not pretrained:
                logger.warning("Prithvi embed mode with a randomly initialised encoder")
            logger.info("Building Prithvi encoder %s (bands %s)", registry_name, names)
            model = BACKBONE_REGISTRY.build(
                registry_name, pretrained=pretrained, bands=PRITHVI_HLS_BANDS, num_frames=1
            )
            model_patch = int(model.patch_embed.patch_size[-1])
            if model_patch != patch:
                raise RuntimeError(
                    f"{registry_name} has patch size {model_patch}, expected {patch} "
                    "(agribound's PRITHVI_ARCH table is out of date for this terratorch)"
                )
            if pretrained:
                hub = hf_cached_file_info(*PRITHVI_WEIGHTS[registry_name])
                meta["weights_revision"] = hub["revision"]
                meta["weights_sha256"] = hub["sha256"]
            model = model.float().eval().to(device)
            with rasterio.open(raster_path) as src:
                height, width, nodata = src.height, src.width, src.nodata
                profile = {
                    "driver": "GTiff",
                    "height": height,
                    "width": width,
                    "count": 1,
                    "dtype": "int32",
                    "crs": src.crs,
                    "transform": src.transform,
                    "nodata": 0,
                    "compress": "lzw",
                    "tiled": True,
                    "blockxsize": 256,
                    "blockysize": 256,
                }

                def read_rows(row0: int, nrows: int) -> tuple[np.ndarray, np.ndarray]:
                    data = src.read(indices, window=Window(0, row0, width, nrows))
                    return self._normalise(data, config.source, value_scale, nodata)

                tokens, valid = extract_token_map(
                    read_rows,
                    height,
                    width,
                    model,
                    tile=tile,
                    patch=patch,
                    batch_size=batch_size,
                    device=device,
                    coords=coords,
                    layer=layer,
                    torch_module=torch,
                )
            del model
            _empty_cache(torch, device)

            rng = get_rng(config, "prithvi", "embed", "kmeans-sample")
            flat_valid = np.flatnonzero(valid)
            if flat_valid.size == 0:
                logger.warning("No valid pixels in %s", raster_path)
                return _empty(profile["crs"], meta)
            n_sample = min(50_000, flat_valid.size)
            pick = np.sort(rng.choice(flat_valid, n_sample, replace=False))
            ys, xs = np.divmod(pick, width)
            sample = interpolate_tokens(tokens, patch, ys, xs)
            km, k, score = fit_kmeans(sample, n_clusters, config.seed)
            meta.update({"n_clusters": int(k), "silhouette": score, "n_fit_samples": int(n_sample)})

            tmp = cluster_path.with_name(cluster_path.stem + ".partial.tif")
            rows_per_strip = max(1, int(_PREDICT_BUDGET_BYTES // (width * tokens.shape[-1] * 4)))
            with rasterio.open(tmp, "w", **profile) as dst:
                for row0 in range(0, height, rows_per_strip):
                    nrows = min(rows_per_strip, height - row0)
                    ys_strip = np.arange(row0, row0 + nrows)
                    emb = interpolate_token_rows(tokens, patch, ys_strip, width)
                    labels = km.predict(emb.reshape(-1, emb.shape[-1])).reshape(nrows, width) + 1
                    labels = np.where(valid[row0 : row0 + nrows], labels, 0).astype(np.int32)
                    dst.write(labels, 1, window=Window(0, row0, width, nrows))
            tmp.replace(cluster_path)
            write_json(
                sidecar,
                {
                    k_: meta.get(k_)
                    for k_ in (
                        "n_clusters",
                        "silhouette",
                        "n_fit_samples",
                        "device",
                        "weights_revision",
                        "weights_sha256",
                    )
                },
            )

        from agribound.postprocess.polygonize import polygonize_mask

        gdf = polygonize_mask(str(cluster_path), min_area_m2=config.min_field_area_m2)
        gdf.attrs["engine_meta"] = meta
        logger.info("Prithvi embedding clustering delineated %d polygons", len(gdf))
        return gdf

    # ------------------------------------------------------------------
    # Segment mode
    # ------------------------------------------------------------------

    def _segment_mode(
        self, raster_path: str, config: AgriboundConfig, checkpoint: str
    ) -> gpd.GeoDataFrame:
        """Prithvi + UPerNet segmentation with a terratorch checkpoint."""
        import rasterio

        from agribound._cache import cache_path
        from agribound.engines.finetune._data import (
            file_sha256,
            interior_polygons,
            package_versions,
            raster_fingerprint,
            read_checkpoint_hparams,
            read_json,
            write_json,
        )

        torch = _import_terratorch("segmentation mode")
        from terratorch.tasks import SemanticSegmentationTask
        from terratorch.tasks.tiled_inference import tiled_inference

        ckpt = Path(checkpoint).expanduser()
        if not ckpt.is_file():
            raise FileNotFoundError(f"Prithvi checkpoint not found: {ckpt}")
        params = config.engine_params
        tile = int(params.get("tile_size", params.get("patch_size", 224)))
        stride = int(params.get("stride", max(tile - 32, 1)))
        batch_size = max(1, int(params.get("batch_size", 8)))
        value_scale = params.get("value_scale")
        names, indices = self._bands(config)
        device = config.resolve_device()
        training = _read_training_meta(ckpt)
        dilate_px = params.get("dilate_interior_px")
        if dilate_px is None:
            dilate_px = training.get("boundary_erosion", params.get("boundary_erosion", 2))
        dilate_px = int(dilate_px)

        pred_path = cache_path(
            config,
            "prithvi_segmentation",
            ".tif",
            raster_fingerprint(raster_path),
            raster_fingerprint(ckpt),
            f"bands={names}:{indices}",
            f"value_scale={value_scale}",
            f"tile={tile}",
            f"stride={stride}",
            "pad=reflect-v2",
        )
        hparams = read_checkpoint_hparams(ckpt)
        model_args = dict(hparams.get("model_args") or {})
        meta: dict[str, Any] = {
            "backend": "terratorch",
            "mode": "segment",
            "checkpoint": str(ckpt.resolve()),
            "checkpoint_sha256": file_sha256(ckpt),
            "model": model_args.get("backbone"),
            "decoder": model_args.get("decoder"),
            "necks": model_args.get("necks"),
            "peft_config": model_args.get("peft_config"),
            "num_classes": model_args.get("num_classes"),
            "band_names": names,
            "band_indices": indices,
            "input_units": "surface reflectance x 10000",
            "normalisation": {"mean": PRITHVI_MEAN, "std": PRITHVI_STD},
            "tile_size": tile,
            "stride": stride,
            "field_class": 1,
            "dilate_interior_px": dilate_px,
            "device": device,
            "training": training or None,
        }
        meta.update(package_versions("terratorch", "torch"))

        pred_info = pred_path.with_suffix(".json")
        if pred_path.exists() and pred_info.exists():
            logger.info("Using cached Prithvi segmentation: %s", pred_path)
            meta.update(read_json(pred_info), cache_reused=True)
        else:
            # Weights come from the checkpoint; do not download the backbone again.
            load_args = {**model_args, "backbone_pretrained": False}
            task = SemanticSegmentationTask.load_from_checkpoint(
                str(ckpt), map_location="cpu", model_args=load_args
            )
            with rasterio.open(raster_path) as src:
                data = src.read(indices)
                profile = {
                    "driver": "GTiff",
                    "height": src.height,
                    "width": src.width,
                    "count": 1,
                    "dtype": "uint8",
                    "crs": src.crs,
                    "transform": src.transform,
                    "compress": "lzw",
                }
                nodata = src.nodata
            norm, valid = self._normalise(data, config.source, value_scale, nodata)
            del data
            height, width = valid.shape
            single_pass = height <= tile and width <= tile
            if single_pass:
                # tiled_inference runs one forward pass on the input as given;
                # the model reflect-pads it to a multiple of 2 x patch.
                model_input = (height, width)
                pad_h = pad_w = 0
            else:
                # tiled_inference needs both sides >= the tile: mirror-pad a
                # shorter side (bottom/right) to the tile size.
                model_input = (tile, tile)
                pad_h, pad_w = max(0, tile - height), max(0, tile - width)
                if pad_h or pad_w:
                    norm = np.pad(norm, ((0, 0), (0, pad_h), (0, pad_w)), mode="reflect")
            meta["model_input_size"] = list(model_input)
            meta["input_padding"] = (
                {"mode": "reflect", "rows": pad_h, "cols": pad_w} if pad_h or pad_w else None
            )
            patch = _patch_size_of(model_args, task.model)
            if patch is not None:
                device = upernet_device(device, model_input, patch, "segmentation")
                meta["device"] = device
            model = task.model.float().eval().to(device)
            x = torch.from_numpy(norm)[None]
            if single_pass:
                x = x.to(device)

            def forward(t):
                return model(t).output

            logits = tiled_inference(
                forward,
                x,
                h_crop=tile,
                w_crop=tile,
                h_stride=stride,
                w_stride=stride,
                batch_size=batch_size,
                device=device,
            )
            pred = logits.argmax(dim=1)[0].cpu().numpy().astype(np.uint8)[:height, :width]
            pred[~valid] = 0
            del model, task, logits
            _empty_cache(torch, device)
            tmp = pred_path.with_name(pred_path.stem + ".partial.tif")
            with rasterio.open(tmp, "w", **profile) as dst:
                dst.write(pred, 1)
            tmp.replace(pred_path)
            write_json(
                pred_info,
                {k: meta[k] for k in ("device", "model_input_size", "input_padding")},
            )

        gdf = interior_polygons(pred_path, dilate_px, config.min_field_area_m2)
        gdf.attrs["engine_meta"] = meta
        logger.info("Prithvi segmentation delineated %d fields", len(gdf))
        return gdf

    # ------------------------------------------------------------------
    # PCA mode
    # ------------------------------------------------------------------

    def _pca_mode(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
        """K-means on PCA of per-band z-scores of R, G, B, NIR (no ViT)."""
        from agribound._cache import cache_path
        from agribound._repro import get_rng
        from agribound.engines.base import get_canonical_band_indices
        from agribound.engines.finetune._data import (
            raster_fingerprint,
            read_json,
            valid_pixels,
            write_json,
        )
        from agribound.io.raster import read_raster, write_raster

        params = config.engine_params
        n_clusters = params.get("n_clusters", "auto")
        indices = get_canonical_band_indices(
            config.source, ["R", "G", "B", "NIR"], bands=config.bands
        )
        cluster_path = cache_path(
            config,
            "prithvi_pca_clusters",
            ".tif",
            raster_fingerprint(raster_path),
            f"bands={indices}",
            f"k={n_clusters}",
            f"seed={config.seed}",
        )
        meta: dict[str, Any] = {
            "backend": "scikit-learn",
            "mode": "pca",
            "band_indices": indices,
            "seed": config.seed,
        }
        sidecar = cluster_path.with_suffix(".json")
        if cluster_path.exists() and sidecar.exists():
            logger.info("Using cached PCA clusters: %s", cluster_path)
            meta.update(read_json(sidecar), cache_reused=True)
        else:
            data, raster_meta = read_raster(raster_path, bands=indices)
            valid = valid_pixels(data, raster_meta.get("nodata"))
            if not valid.any():
                logger.warning("No valid pixels in %s", raster_path)
                return _empty(raster_meta.get("crs"), meta)
            features = pca_features(data, valid, config.seed)
            rng = get_rng(config, "prithvi", "pca", "kmeans-sample")
            n_valid = features.shape[0]
            pick = rng.choice(n_valid, min(50_000, n_valid), replace=False)
            km, k, score = fit_kmeans(features[np.sort(pick)], n_clusters, config.seed)
            labels = np.zeros(valid.shape, dtype=np.int32)
            labels[valid] = km.predict(features) + 1
            write_raster(
                cluster_path,
                labels[np.newaxis],
                crs=raster_meta["crs"],
                transform=raster_meta["transform"],
                nodata=0,
            )
            meta.update({"n_clusters": int(k), "silhouette": score})
            write_json(sidecar, {"n_clusters": int(k), "silhouette": score})

        from agribound.postprocess.polygonize import polygonize_mask

        gdf = polygonize_mask(str(cluster_path), min_area_m2=config.min_field_area_m2)
        gdf.attrs["engine_meta"] = meta
        logger.info("PCA clustering delineated %d polygons", len(gdf))
        return gdf

    # ------------------------------------------------------------------
    # Prefetch
    # ------------------------------------------------------------------

    @classmethod
    def prefetch(cls, config: AgriboundConfig) -> list[str]:
        """Download the Prithvi-EO-2.0 weights of ``engine_params["model_name"]``.

        The pre-trained weights are needed by ``embed`` mode (unless
        ``engine_params["pretrained"]`` is *False*) and by fine-tuning
        (``config.fine_tune``, unless ``backbone_pretrained`` is *False*).
        ``segment`` inference loads every weight from its checkpoint and
        ``pca`` mode uses no model, so nothing is downloaded for them. Files
        go to the Hugging Face cache (``HF_HOME``); set ``HF_HUB_OFFLINE=1``
        on nodes without internet access afterwards.

        Returns
        -------
        list[str]
            Local path of the weights file when needed, plus the checkpoint
            if ``engine_params["checkpoint_path"]`` exists.
        """
        from huggingface_hub import hf_hub_download

        params = config.engine_params
        checkpoint = params.get("checkpoint_path")
        mode = str(params.get("mode") or ("segment" if checkpoint else "embed")).lower()
        needs_weights = (config.fine_tune and params.get("backbone_pretrained", True)) or (
            mode == "embed" and params.get("pretrained", True)
        )
        paths: list[str] = []
        if needs_weights:
            registry_name = resolve_prithvi_model(params.get("model_name"))
            repo, filename = PRITHVI_WEIGHTS[registry_name]
            paths.append(hf_hub_download(repo_id=repo, filename=filename))
        else:
            logger.info("Prithvi mode=%r needs no pre-trained weights; nothing downloaded", mode)
        if checkpoint and Path(checkpoint).is_file():
            paths.append(str(Path(checkpoint).resolve()))
        return paths
delineate
delineate(raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame

Run Prithvi-based field delineation.

Parameters:

Name Type Description Default
raster_path str

Composite GeoTIFF.

required
config AgriboundConfig

Pipeline configuration. engine_params:

  • mode: "embed" | "segment" | "pca" (default: see the module docstring).
  • checkpoint_path: terratorch SemanticSegmentationTask checkpoint (segment mode).
  • model_name: Prithvi-EO-2.0 variant for embed mode (default "Prithvi-EO-2.0-300M-TL"; see :data:PRITHVI_MODELS).
  • tile_size: tile edge in pixels (default 224). In embed mode it must be a multiple of the patch size (16; 14 for the 600M models); tiles that extend past the raster are filled by mirror reflection of the raster. In segment mode a raster that fits in one tile is passed whole (terratorch reflect-pads it to a multiple of twice the patch size); a larger raster is run through terratorch's tiled_inference, after mirror padding a side shorter than the tile to the tile size. On Apple MPS the segmentation model runs on CPU (logged at WARNING) unless the model input size passes :func:upernet_mps_compatible (e.g. 192 px tiles for patch 16). patch_size is accepted as a legacy alias.
  • stride: tile step for segment mode (default tile_size - 32).
  • batch_size: tiles per forward pass (default 8).
  • n_clusters: int or "auto" (silhouette over 5, 10, 15, 20, 30 on a seeded sample) for embed/pca.
  • embed_layer: encoder layer used in embed mode (default -1, the normalised last layer).
  • temporal_coords: [year, day_of_year], or False to not pass them (default: :func:composite_mid_date); embed mode with a *-TL model only (segment mode and fine-tuning pass no coordinates).
  • location_coords: [lat, lon], or False (default: the raster centre); embed mode with a *-TL model only.
  • dilate_interior_px: segment mode; pixels by which each interior region is grown over the predicted boundary class (default: the checkpoint's training boundary_erosion if recorded, else engine_params["boundary_erosion"] or 2; see :func:agribound.engines.finetune._data.interior_polygons).
  • value_scale: required for source="local" ("reflectance_x10000" or "unit"). Without config.bands a local raster's bands 1-6 are read as Blue, Green, Red, narrow NIR, SWIR 1, SWIR 2 (pca mode: bands 1-4 as R, G, B, NIR).
  • pretrained: False builds a randomly initialised encoder in embed mode (for tests only; logged as a warning).
required

Returns:

Type Description
GeoDataFrame

Polygons with gdf.attrs["engine_meta"].

Raises:

Type Description
RuntimeError

If mode="segment" is requested without a checkpoint.

ValueError

For unknown modes, models or non-reflectance inputs.

Source code in agribound/engines/prithvi.py
def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
    """Run Prithvi-based field delineation.

    Parameters
    ----------
    raster_path : str
        Composite GeoTIFF.
    config : AgriboundConfig
        Pipeline configuration. ``engine_params``:

        - ``mode``: ``"embed"`` | ``"segment"`` | ``"pca"`` (default: see
          the module docstring).
        - ``checkpoint_path``: terratorch ``SemanticSegmentationTask``
          checkpoint (``segment`` mode).
        - ``model_name``: Prithvi-EO-2.0 variant for ``embed`` mode
          (default ``"Prithvi-EO-2.0-300M-TL"``; see :data:`PRITHVI_MODELS`).
        - ``tile_size``: tile edge in pixels (default 224). In ``embed``
          mode it must be a multiple of the patch size (16; 14 for the
          600M models); tiles that extend past the raster are filled by
          mirror reflection of the raster. In ``segment`` mode a raster
          that fits in one tile is passed whole (terratorch reflect-pads
          it to a multiple of twice the patch size); a larger raster is
          run through terratorch's ``tiled_inference``, after mirror
          padding a side shorter than the tile to the tile size. On
          Apple MPS the segmentation model runs on CPU (logged at
          WARNING) unless the model input size passes
          :func:`upernet_mps_compatible` (e.g. 192 px tiles for patch
          16). ``patch_size`` is accepted as a legacy alias.
        - ``stride``: tile step for ``segment`` mode (default
          ``tile_size - 32``).
        - ``batch_size``: tiles per forward pass (default 8).
        - ``n_clusters``: int or ``"auto"`` (silhouette over 5, 10, 15,
          20, 30 on a seeded sample) for ``embed``/``pca``.
        - ``embed_layer``: encoder layer used in ``embed`` mode (default
          -1, the normalised last layer).
        - ``temporal_coords``: ``[year, day_of_year]``, or *False* to not
          pass them (default: :func:`composite_mid_date`); ``embed``
          mode with a ``*-TL`` model only (``segment`` mode and
          fine-tuning pass no coordinates).
        - ``location_coords``: ``[lat, lon]``, or *False* (default: the
          raster centre); ``embed`` mode with a ``*-TL`` model only.
        - ``dilate_interior_px``: ``segment`` mode; pixels by which each
          interior region is grown over the predicted boundary class
          (default: the checkpoint's training ``boundary_erosion`` if
          recorded, else ``engine_params["boundary_erosion"]`` or 2; see
          :func:`agribound.engines.finetune._data.interior_polygons`).
        - ``value_scale``: required for ``source="local"``
          (``"reflectance_x10000"`` or ``"unit"``). Without
          ``config.bands`` a local raster's bands 1-6 are read as Blue,
          Green, Red, narrow NIR, SWIR 1, SWIR 2 (``pca`` mode: bands
          1-4 as R, G, B, NIR).
        - ``pretrained``: *False* builds a randomly initialised encoder in
          ``embed`` mode (for tests only; logged as a warning).

    Returns
    -------
    geopandas.GeoDataFrame
        Polygons with ``gdf.attrs["engine_meta"]``.

    Raises
    ------
    RuntimeError
        If ``mode="segment"`` is requested without a checkpoint.
    ValueError
        For unknown modes, models or non-reflectance inputs.
    """
    params = config.engine_params
    checkpoint = params.get("checkpoint_path")
    mode = str(params.get("mode") or ("segment" if checkpoint else "embed")).lower()
    if mode not in _MODES:
        raise ValueError(f"Unknown Prithvi mode {mode!r}. Choose from {_MODES}")
    self.validate_input(raster_path, config)
    if mode == "segment":
        if not checkpoint:
            raise RuntimeError(
                "Prithvi mode='segment' needs a fine-tuned checkpoint: set fine_tune=True "
                "with reference_boundaries, or engine_params['checkpoint_path']. Use "
                "mode='embed' for label-free clustering."
            )
        return self._segment_mode(raster_path, config, str(checkpoint))
    if checkpoint:
        logger.warning(
            "Prithvi mode=%r ignores engine_params['checkpoint_path'] (%s)", mode, checkpoint
        )
    if mode == "pca":
        return self._pca_mode(raster_path, config)
    return self._embed_mode(raster_path, config)
prefetch classmethod
prefetch(config: AgriboundConfig) -> list[str]

Download the Prithvi-EO-2.0 weights of engine_params["model_name"].

The pre-trained weights are needed by embed mode (unless engine_params["pretrained"] is False) and by fine-tuning (config.fine_tune, unless backbone_pretrained is False). segment inference loads every weight from its checkpoint and pca mode uses no model, so nothing is downloaded for them. Files go to the Hugging Face cache (HF_HOME); set HF_HUB_OFFLINE=1 on nodes without internet access afterwards.

Returns:

Type Description
list[str]

Local path of the weights file when needed, plus the checkpoint if engine_params["checkpoint_path"] exists.

Source code in agribound/engines/prithvi.py
@classmethod
def prefetch(cls, config: AgriboundConfig) -> list[str]:
    """Download the Prithvi-EO-2.0 weights of ``engine_params["model_name"]``.

    The pre-trained weights are needed by ``embed`` mode (unless
    ``engine_params["pretrained"]`` is *False*) and by fine-tuning
    (``config.fine_tune``, unless ``backbone_pretrained`` is *False*).
    ``segment`` inference loads every weight from its checkpoint and
    ``pca`` mode uses no model, so nothing is downloaded for them. Files
    go to the Hugging Face cache (``HF_HOME``); set ``HF_HUB_OFFLINE=1``
    on nodes without internet access afterwards.

    Returns
    -------
    list[str]
        Local path of the weights file when needed, plus the checkpoint
        if ``engine_params["checkpoint_path"]`` exists.
    """
    from huggingface_hub import hf_hub_download

    params = config.engine_params
    checkpoint = params.get("checkpoint_path")
    mode = str(params.get("mode") or ("segment" if checkpoint else "embed")).lower()
    needs_weights = (config.fine_tune and params.get("backbone_pretrained", True)) or (
        mode == "embed" and params.get("pretrained", True)
    )
    paths: list[str] = []
    if needs_weights:
        registry_name = resolve_prithvi_model(params.get("model_name"))
        repo, filename = PRITHVI_WEIGHTS[registry_name]
        paths.append(hf_hub_download(repo_id=repo, filename=filename))
    else:
        logger.info("Prithvi mode=%r needs no pre-trained weights; nothing downloaded", mode)
    if checkpoint and Path(checkpoint).is_file():
        paths.append(str(Path(checkpoint).resolve()))
    return paths

Embedding clustering

embedding

Embedding-based clustering engine.

Clusters pre-computed per-pixel embeddings (Google Satellite Embedding V1, 64-D, or TESSERA, 128-D) with scikit-learn and polygonizes the cluster map. No labels, model weights or GPU are needed. Clusters are land-cover segments, not field instances: every connected region of every cluster becomes a polygon, and non-cropland segments are only removed by the downstream area and LULC filters.

Clustering is memory-bounded: the raster is read in row blocks of at most engine_params["max_block_mb"] MiB, a seeded uniform random sample of valid pixels is drawn in one pass, the dimensionality reduction and the clusterer are fitted on that sample, and a second pass predicts the label of every valid pixel block by block into an int32 label raster. Polygonization (:func:agribound.postprocess.polygonize.polygonize_mask) then reads that label raster at once (4 bytes per pixel).

EmbeddingEngine

Bases: DelineationEngine

Field delineation by unsupervised clustering of pixel embeddings.

Engine parameters (config.engine_params)

use_pca : bool Reduce the embeddings with PCA before clustering (default True; only when the raster has more bands than pca_components). pca_components : int PCA dimensions (default 16). matryoshka_depth : int or None TESSERA v2 only (source="tessera-embedding" with tessera_version="v2"): cluster the first depth dimensions, a Matryoshka prefix, instead of PCA. One of :data:MATRYOSHKA_DEPTHS. Any other source or version raises ValueError. n_clusters : int or "auto" Number of clusters, or "auto" (default): the candidate in k_candidates with the highest silhouette score of a KMeans(n_init=10) fit on the first silhouette_sample_size sampled pixels. clustering_method : str "kmeans" (default) or "spectral". "kmeans" fits KMeans(n_init=10) on the cluster sample whatever the raster size and keeps the lowest-inertia of the ten complete restarts (0.1.x and 1.0.0 used MiniBatchKMeans(batch_size=10000, n_init=3) above 100 000 valid pixels, which ends in a higher-inertia solution for many samples, and KMeans(n_init=5) otherwise). "spectral" is SpectralClustering(affinity="nearest_neighbors") on the sample, extended to all pixels with NearestCentroid (slow). k_candidates : list[int] Candidates for n_clusters="auto" (default 5, 10, 15, 20, 30, 50). pca_sample_size, cluster_sample_size, silhouette_sample_size : int Sizes of the nested random samples used to fit PCA, the clusterer and the silhouette selection (defaults 100 000, 50 000, 5 000). max_block_mb : float Maximum size of one row block read from the raster (default 256 MiB). Peak memory is a few times this (a block, the copy of its valid pixels and per-band masks) plus the fitting sample; GDAL's block cache is limited to max(64, max_block_mb) MiB during the reads unless GDAL_CACHEMAX is set in the environment.

A pixel is valid when all bands used (every band, or the Matryoshka prefix) are finite, not all zero and, when the raster declares a finite nodata value, not all equal to it. All random choices (pixel sample, PCA solver, k-means initialisation) are seeded from config.seed. scikit-learn computes the k-means sums in parallel, so a different number of OpenMP threads can move a small fraction of pixels to another cluster. The cluster raster (int32; 0 = invalid, 1..k = cluster) is cached with :func:agribound._cache.cache_path, keyed by the study area, source, year, TESSERA version, every parameter above except max_block_mb (the result does not depend on the block size), the seed and the input raster's path, size and modification time.

When config.sam_refine is True (which also absorbs the legacy engine_params["sam_refine"]), the polygons are refined with :func:agribound.engines.samgeo_engine.refine_boundaries on this raster. An embedding raster has no RGB bands, so this requires engine_params["sam_rgb_bands"] (three 1-based dimensions used as a pseudo-RGB image); without it the engine raises before clustering.

Source code in agribound/engines/embedding.py
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
class EmbeddingEngine(DelineationEngine):
    """Field delineation by unsupervised clustering of pixel embeddings.

    Engine parameters (``config.engine_params``)
    --------------------------------------------
    use_pca : bool
        Reduce the embeddings with PCA before clustering (default *True*;
        only when the raster has more bands than *pca_components*).
    pca_components : int
        PCA dimensions (default 16).
    matryoshka_depth : int or None
        TESSERA v2 only (``source="tessera-embedding"`` with
        ``tessera_version="v2"``): cluster the first *depth* dimensions,
        a Matryoshka prefix, instead of PCA. One of :data:`MATRYOSHKA_DEPTHS`.
        Any other source or version raises ``ValueError``.
    n_clusters : int or "auto"
        Number of clusters, or ``"auto"`` (default): the candidate in
        *k_candidates* with the highest silhouette score of a
        ``KMeans(n_init=10)`` fit on the first *silhouette_sample_size*
        sampled pixels.
    clustering_method : str
        ``"kmeans"`` (default) or ``"spectral"``. ``"kmeans"`` fits
        ``KMeans(n_init=10)`` on the cluster sample whatever the raster size
        and keeps the lowest-inertia of the ten complete restarts (0.1.x and
        1.0.0 used ``MiniBatchKMeans(batch_size=10000, n_init=3)`` above
        100 000 valid pixels, which ends in a higher-inertia solution for
        many samples, and ``KMeans(n_init=5)`` otherwise). ``"spectral"`` is
        ``SpectralClustering(affinity="nearest_neighbors")`` on the sample,
        extended to all pixels with ``NearestCentroid`` (slow).
    k_candidates : list[int]
        Candidates for ``n_clusters="auto"`` (default 5, 10, 15, 20, 30, 50).
    pca_sample_size, cluster_sample_size, silhouette_sample_size : int
        Sizes of the nested random samples used to fit PCA, the clusterer and
        the silhouette selection (defaults 100 000, 50 000, 5 000).
    max_block_mb : float
        Maximum size of one row block read from the raster (default 256 MiB).
        Peak memory is a few times this (a block, the copy of its valid
        pixels and per-band masks) plus the fitting sample; GDAL's block
        cache is limited to ``max(64, max_block_mb)`` MiB during the reads
        unless ``GDAL_CACHEMAX`` is set in the environment.

    A pixel is valid when all bands used (every band, or the Matryoshka
    prefix) are finite, not all zero and, when the raster declares a finite
    nodata value, not all equal to it. All
    random choices (pixel sample, PCA solver, k-means initialisation) are
    seeded from ``config.seed``. scikit-learn computes the k-means sums in
    parallel, so a different number of OpenMP threads can move a small
    fraction of pixels to another cluster. The cluster raster (int32;
    0 = invalid, ``1..k`` = cluster) is cached with :func:`agribound._cache.cache_path`,
    keyed by the study area, source, year, TESSERA version, every parameter
    above except *max_block_mb* (the result does not depend on the block
    size), the seed and the input raster's path, size and modification time.

    When ``config.sam_refine`` is *True* (which also absorbs the legacy
    ``engine_params["sam_refine"]``), the polygons are refined with
    :func:`agribound.engines.samgeo_engine.refine_boundaries` on this raster.
    An embedding raster has no RGB bands, so this requires
    ``engine_params["sam_rgb_bands"]`` (three 1-based dimensions used as a
    pseudo-RGB image); without it the engine raises before clustering.
    """

    name = "embedding"
    supported_sources = list(ENGINE_REGISTRY["embedding"]["supported_sources"])
    requires_bands = list(ENGINE_REGISTRY["embedding"]["requires_bands"])

    # ------------------------------------------------------------------
    # Parameters
    # ------------------------------------------------------------------

    @staticmethod
    def resolve_params(config: AgriboundConfig) -> dict[str, Any]:
        """Return the validated engine parameters (defaults filled in).

        Raises
        ------
        ValueError
            For invalid values.
        """
        user = dict(config.engine_params or {})
        params = {k: user.get(k, v) for k, v in DEFAULT_PARAMS.items()}
        method = str(params["clustering_method"]).lower().strip()
        if method not in ("kmeans", "spectral"):
            raise ValueError(f"clustering_method must be 'kmeans' or 'spectral', got {method!r}")
        params["clustering_method"] = method
        n_clusters = params["n_clusters"]
        if isinstance(n_clusters, str):
            if n_clusters.lower().strip() != "auto":
                raise ValueError(
                    f"n_clusters must be an integer >= 2 or 'auto', got {n_clusters!r}"
                )
            params["n_clusters"] = "auto"
        elif isinstance(n_clusters, bool) or int(n_clusters) != n_clusters or n_clusters < 2:
            raise ValueError(f"n_clusters must be an integer >= 2 or 'auto', got {n_clusters!r}")
        else:
            params["n_clusters"] = int(n_clusters)
        candidates = [int(k) for k in params["k_candidates"]]
        if not candidates or min(candidates) < 2:
            raise ValueError(f"k_candidates must be integers >= 2, got {params['k_candidates']!r}")
        params["k_candidates"] = sorted(set(candidates))
        for key in ("pca_components", "pca_sample_size", "cluster_sample_size"):
            params[key] = int(params[key])
            if params[key] < 1:
                raise ValueError(f"{key} must be >= 1, got {params[key]}")
        params["silhouette_sample_size"] = int(params["silhouette_sample_size"])
        if params["silhouette_sample_size"] < 3:
            raise ValueError("silhouette_sample_size must be >= 3")
        params["use_pca"] = bool(params["use_pca"])
        params["max_block_mb"] = float(params["max_block_mb"])
        if params["max_block_mb"] <= 0:
            raise ValueError("max_block_mb must be > 0")

        depth = params["matryoshka_depth"]
        if depth is not None:
            depth = int(depth)
            if config.source != "tessera-embedding" or config.tessera_version != "v2":
                raise ValueError(
                    "matryoshka_depth needs TESSERA v2 embeddings (source='tessera-embedding', "
                    f"tessera_version='v2'); got source={config.source!r}, "
                    f"tessera_version={config.tessera_version!r}. Remove it to use PCA."
                )
            if depth not in MATRYOSHKA_DEPTHS:
                raise ValueError(
                    f"matryoshka_depth must be one of {MATRYOSHKA_DEPTHS}, got {depth}"
                )
            params["matryoshka_depth"] = depth
        return params

    # ------------------------------------------------------------------
    # Engine API
    # ------------------------------------------------------------------

    @staticmethod
    def cluster_cache_path(
        raster_path: str, config: AgriboundConfig, params: dict[str, Any] | None = None
    ) -> Path:
        """Return the cache path of the cluster raster for *raster_path* and *config*.

        The key covers :func:`agribound._cache.cache_key`'s fields plus the
        seed, the raster's resolved path, size and modification time, and all
        engine parameters except ``max_block_mb``.
        """
        from agribound._cache import cache_path

        params = params if params is not None else EmbeddingEngine.resolve_params(config)
        # max_block_mb only changes how the raster is read, not the result.
        parts = [_CACHE_VERSION, f"seed={config.seed}", _file_signature(raster_path)] + [
            f"{k}={params[k]}" for k in sorted(params) if k != "max_block_mb"
        ]
        return cache_path(config, "embedding_clusters", ".tif", *parts)

    @classmethod
    def prefetch(cls, config: AgriboundConfig) -> list[str]:
        """Download the SAM weights when ``config.sam_refine`` is set (nothing else is remote)."""
        if not config.sam_refine:
            return []
        from agribound.engines.samgeo_engine import prefetch as sam_prefetch

        return sam_prefetch(config)

    def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
        """Cluster the embedding raster and polygonize the clusters.

        Parameters
        ----------
        raster_path : str
            Embedding GeoTIFF (float32, one band per embedding dimension).
        config : AgriboundConfig
            Pipeline configuration.

        Returns
        -------
        geopandas.GeoDataFrame
            Polygons with ``class_value`` (cluster label ``1..k``), in the
            raster CRS. ``attrs["engine_meta"]`` describes the clustering;
            ``attrs["sam_stats"]`` is set when SAM refinement ran.

        Raises
        ------
        ValueError
            For an unsupported source, invalid parameters, a raster with too
            few valid pixels, or SAM refinement without ``sam_rgb_bands``.
        """
        from agribound.io.raster import get_raster_info

        if config.source not in self.supported_sources:
            raise ValueError(
                f"The embedding engine needs an embedding source {self.supported_sources}, "
                f"got {config.source!r}"
            )
        params = self.resolve_params(config)
        info = get_raster_info(raster_path)
        depth = params["matryoshka_depth"]
        if depth is not None and depth > info.count:
            raise ValueError(f"matryoshka_depth={depth} but the raster has only {info.count} bands")
        if config.sam_refine:
            from agribound.engines.samgeo_engine import _rgb_band_indices

            _rgb_band_indices(config, info.count)  # fail before clustering

        cluster_path = self.cluster_cache_path(raster_path, config, params)
        meta_path = cluster_path.with_suffix(".json")

        if cluster_path.exists() and meta_path.exists():
            logger.info("Using cached embedding clusters: %s", cluster_path)
            meta = json.loads(meta_path.read_text())
            meta["cache_hit"] = True
        else:
            logger.info(
                "Clustering embeddings: %d bands, %dx%d pixels", info.count, info.width, info.height
            )
            meta = self._cluster_raster(raster_path, cluster_path, params, config)
            _atomic_write_text(meta_path, json.dumps(meta, indent=2, sort_keys=True))
            meta["cache_hit"] = False
        meta["cluster_raster"] = str(cluster_path)

        from agribound.postprocess.polygonize import polygonize_mask

        gdf = polygonize_mask(str(cluster_path), min_area_m2=config.min_field_area_m2)
        logger.info("Embedding clustering delineated %d polygons", len(gdf))

        meta["sam_refine"] = bool(config.sam_refine)
        if config.sam_refine and len(gdf) > 0:
            from agribound.engines.samgeo_engine import refine_boundaries

            gdf = refine_boundaries(gdf, raster_path, config)
            meta["sam_stats"] = gdf.attrs.get("sam_stats")
        gdf.attrs["engine_meta"] = meta
        return gdf

    # ------------------------------------------------------------------
    # Clustering
    # ------------------------------------------------------------------

    def _cluster_raster(
        self,
        raster_path: str,
        out_path: Path,
        params: dict[str, Any],
        config: AgriboundConfig,
    ) -> dict[str, Any]:
        """Fit on a seeded sample, predict block by block, write the int32 label raster."""
        import rasterio
        import sklearn
        from rasterio.windows import Window

        from agribound._repro import get_rng

        seed = int(config.seed)
        depth = params["matryoshka_depth"]
        # Bound GDAL's block cache too (default: 5 % of RAM) unless the user set it.
        env = {} if "GDAL_CACHEMAX" in os.environ else {"GDAL_CACHEMAX": _gdal_cache_mb(params)}
        with rasterio.Env(**env), rasterio.open(raster_path) as src:
            n_bands = src.count
            bands = list(range(1, (depth or n_bands) + 1))
            height, width = src.height, src.width
            nodata = src.nodata
            bytes_per_row = max(1, width * len(bands) * 4)
            block_rows = int(max(1, min(height, params["max_block_mb"] * 2**20 // bytes_per_row)))

            def blocks():
                for row in range(0, height, block_rows):
                    h = min(block_rows, height - row)
                    data = src.read(bands, window=Window(0, row, width, h)).astype(
                        np.float32, copy=False
                    )
                    flat = data.reshape(len(bands), -1).T
                    yield row, h, flat, _valid_mask(data, nodata)

            # Pass 1: uniform random sample without replacement (bottom-k random keys).
            n_keep = max(params["pca_sample_size"], params["cluster_sample_size"])
            rng = get_rng(config, "embedding-sample")
            keys = np.empty(0, dtype=np.float64)
            sample = np.empty((0, len(bands)), dtype=np.float32)
            n_valid = 0
            for _, _, flat, valid in blocks():
                vals = flat[valid]
                n_valid += len(vals)
                block_keys = rng.random(len(vals))
                if len(block_keys) > n_keep:
                    sel = np.argpartition(block_keys, n_keep - 1)[:n_keep]
                    block_keys, vals = block_keys[sel], vals[sel]
                keys = np.concatenate([keys, block_keys])
                sample = np.concatenate([sample, vals])
                if len(keys) > n_keep:
                    sel = np.argpartition(keys, n_keep - 1)[:n_keep]
                    keys, sample = keys[sel], sample[sel]
            sample = sample[np.argsort(keys, kind="stable")]
            if n_valid < 3:
                raise ValueError(
                    f"{raster_path} has {n_valid} valid embedding pixels; at least 3 are needed"
                )

            # Dimensionality reduction.
            reducer = None
            meta: dict[str, Any] = {
                "backend": "scikit-learn",
                "sklearn_version": sklearn.__version__,
                "source": config.source,
                "tessera_version": config.tessera_version
                if config.source == "tessera-embedding"
                else None,
                "n_bands": int(n_bands),
                "n_pixels": int(height * width),
                "n_valid_pixels": int(n_valid),
                "seed": seed,
                "block_rows": block_rows,
                "params": _jsonable(params),
            }
            if depth is not None:
                meta["reduction"] = "matryoshka"
                meta["n_features"] = depth
            elif params["use_pca"] and n_bands > params["pca_components"]:
                from sklearn.decomposition import PCA

                pca_sample = sample[: params["pca_sample_size"]]
                reducer = PCA(n_components=params["pca_components"], random_state=seed)
                reducer.fit(pca_sample)
                meta["reduction"] = "pca"
                meta["n_features"] = params["pca_components"]
                meta["pca_sample_size"] = int(len(pca_sample))
                meta["pca_explained_variance_ratio"] = float(
                    np.sum(reducer.explained_variance_ratio_)
                )
            else:
                meta["reduction"] = "none"
                meta["n_features"] = int(n_bands)

            fit_x = sample[: params["cluster_sample_size"]]
            if reducer is not None:
                fit_x = reducer.transform(fit_x)
            meta["cluster_sample_size"] = int(len(fit_x))

            k = params["n_clusters"]
            if k == "auto":
                k, scores, n_eval = _select_k(
                    fit_x[: params["silhouette_sample_size"]], params["k_candidates"], seed
                )
                meta["auto_k"] = {
                    "candidates": params["k_candidates"],
                    "silhouette": scores,
                    "sample_size": n_eval,
                }
            if k > len(fit_x):
                raise ValueError(f"n_clusters={k} exceeds the {len(fit_x)} sampled pixels")
            meta["n_clusters"] = int(k)

            predict = self._fit_clusterer(fit_x, k, params["clustering_method"], seed, meta)
            logger.info(
                "Clustering %d valid pixels into %d clusters (%s, reduction=%s)",
                n_valid,
                k,
                meta["clusterer"],
                meta["reduction"],
            )

            # Pass 2: predict per block and write the label raster.
            profile = {
                "driver": "GTiff",
                "height": height,
                "width": width,
                "count": 1,
                "dtype": "int32",
                "crs": src.crs,
                "transform": src.transform,
                "nodata": 0,
                "compress": "lzw",
            }
            tmp_path = out_path.with_name(out_path.stem + ".partial.tif")
            with rasterio.open(tmp_path, "w", **profile) as dst:
                for row, h, flat, valid in blocks():
                    labels = np.zeros(len(flat), dtype=np.int32)
                    idx = np.flatnonzero(valid)
                    for start in range(0, len(idx), _PREDICT_CHUNK):
                        sel = idx[start : start + _PREDICT_CHUNK]
                        x = flat[sel]
                        if reducer is not None:
                            x = reducer.transform(x)
                        labels[sel] = predict(x).astype(np.int32) + 1
                    dst.write(labels.reshape(1, h, width), window=Window(0, row, width, h))
            os.replace(tmp_path, out_path)
        return meta

    @staticmethod
    def _fit_clusterer(x: np.ndarray, k: int, method: str, seed: int, meta: dict[str, Any]):
        """Fit the clusterer on *x* and return a ``predict(array) -> labels`` callable.

        ``"kmeans"`` is ``KMeans(n_init=10)`` on *x* for every raster size; *x*
        has at most *cluster_sample_size* rows, so the ten complete restarts
        stay cheap, and unlike ``MiniBatchKMeans`` they reach the
        lowest-inertia solution for most samples.
        """
        if method == "kmeans":
            from sklearn.cluster import KMeans

            model = KMeans(n_clusters=k, n_init=_KMEANS_N_INIT, random_state=seed)
            model.fit(x)
            meta["clusterer"] = "KMeans"
            meta["kmeans_n_init"] = _KMEANS_N_INIT
            meta["inertia"] = float(model.inertia_)
            return model.predict

        from sklearn.cluster import SpectralClustering
        from sklearn.neighbors import NearestCentroid

        if len(x) > 10_000:
            logger.warning("Spectral clustering with %d samples is slow; consider kmeans", len(x))
        sc = SpectralClustering(n_clusters=k, random_state=seed, affinity="nearest_neighbors")
        sc.fit(x)
        nc = NearestCentroid()
        nc.fit(x, sc.labels_)
        meta["clusterer"] = "SpectralClustering+NearestCentroid"
        return nc.predict

    @staticmethod
    def _auto_select_k(
        sample: np.ndarray, k_range: list[int] | None = None, random_state: int = 42
    ) -> int:
        """Return the silhouette-best k (see :func:`_select_k`)."""
        k_range = k_range if k_range is not None else list(DEFAULT_PARAMS["k_candidates"])
        return _select_k(sample, k_range, random_state)[0]
resolve_params staticmethod
resolve_params(config: AgriboundConfig) -> dict[str, Any]

Return the validated engine parameters (defaults filled in).

Raises:

Type Description
ValueError

For invalid values.

Source code in agribound/engines/embedding.py
@staticmethod
def resolve_params(config: AgriboundConfig) -> dict[str, Any]:
    """Return the validated engine parameters (defaults filled in).

    Raises
    ------
    ValueError
        For invalid values.
    """
    user = dict(config.engine_params or {})
    params = {k: user.get(k, v) for k, v in DEFAULT_PARAMS.items()}
    method = str(params["clustering_method"]).lower().strip()
    if method not in ("kmeans", "spectral"):
        raise ValueError(f"clustering_method must be 'kmeans' or 'spectral', got {method!r}")
    params["clustering_method"] = method
    n_clusters = params["n_clusters"]
    if isinstance(n_clusters, str):
        if n_clusters.lower().strip() != "auto":
            raise ValueError(
                f"n_clusters must be an integer >= 2 or 'auto', got {n_clusters!r}"
            )
        params["n_clusters"] = "auto"
    elif isinstance(n_clusters, bool) or int(n_clusters) != n_clusters or n_clusters < 2:
        raise ValueError(f"n_clusters must be an integer >= 2 or 'auto', got {n_clusters!r}")
    else:
        params["n_clusters"] = int(n_clusters)
    candidates = [int(k) for k in params["k_candidates"]]
    if not candidates or min(candidates) < 2:
        raise ValueError(f"k_candidates must be integers >= 2, got {params['k_candidates']!r}")
    params["k_candidates"] = sorted(set(candidates))
    for key in ("pca_components", "pca_sample_size", "cluster_sample_size"):
        params[key] = int(params[key])
        if params[key] < 1:
            raise ValueError(f"{key} must be >= 1, got {params[key]}")
    params["silhouette_sample_size"] = int(params["silhouette_sample_size"])
    if params["silhouette_sample_size"] < 3:
        raise ValueError("silhouette_sample_size must be >= 3")
    params["use_pca"] = bool(params["use_pca"])
    params["max_block_mb"] = float(params["max_block_mb"])
    if params["max_block_mb"] <= 0:
        raise ValueError("max_block_mb must be > 0")

    depth = params["matryoshka_depth"]
    if depth is not None:
        depth = int(depth)
        if config.source != "tessera-embedding" or config.tessera_version != "v2":
            raise ValueError(
                "matryoshka_depth needs TESSERA v2 embeddings (source='tessera-embedding', "
                f"tessera_version='v2'); got source={config.source!r}, "
                f"tessera_version={config.tessera_version!r}. Remove it to use PCA."
            )
        if depth not in MATRYOSHKA_DEPTHS:
            raise ValueError(
                f"matryoshka_depth must be one of {MATRYOSHKA_DEPTHS}, got {depth}"
            )
        params["matryoshka_depth"] = depth
    return params
cluster_cache_path staticmethod
cluster_cache_path(raster_path: str, config: AgriboundConfig, params: dict[str, Any] | None = None) -> Path

Return the cache path of the cluster raster for raster_path and config.

The key covers :func:agribound._cache.cache_key's fields plus the seed, the raster's resolved path, size and modification time, and all engine parameters except max_block_mb.

Source code in agribound/engines/embedding.py
@staticmethod
def cluster_cache_path(
    raster_path: str, config: AgriboundConfig, params: dict[str, Any] | None = None
) -> Path:
    """Return the cache path of the cluster raster for *raster_path* and *config*.

    The key covers :func:`agribound._cache.cache_key`'s fields plus the
    seed, the raster's resolved path, size and modification time, and all
    engine parameters except ``max_block_mb``.
    """
    from agribound._cache import cache_path

    params = params if params is not None else EmbeddingEngine.resolve_params(config)
    # max_block_mb only changes how the raster is read, not the result.
    parts = [_CACHE_VERSION, f"seed={config.seed}", _file_signature(raster_path)] + [
        f"{k}={params[k]}" for k in sorted(params) if k != "max_block_mb"
    ]
    return cache_path(config, "embedding_clusters", ".tif", *parts)
prefetch classmethod
prefetch(config: AgriboundConfig) -> list[str]

Download the SAM weights when config.sam_refine is set (nothing else is remote).

Source code in agribound/engines/embedding.py
@classmethod
def prefetch(cls, config: AgriboundConfig) -> list[str]:
    """Download the SAM weights when ``config.sam_refine`` is set (nothing else is remote)."""
    if not config.sam_refine:
        return []
    from agribound.engines.samgeo_engine import prefetch as sam_prefetch

    return sam_prefetch(config)
delineate
delineate(raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame

Cluster the embedding raster and polygonize the clusters.

Parameters:

Name Type Description Default
raster_path str

Embedding GeoTIFF (float32, one band per embedding dimension).

required
config AgriboundConfig

Pipeline configuration.

required

Returns:

Type Description
GeoDataFrame

Polygons with class_value (cluster label 1..k), in the raster CRS. attrs["engine_meta"] describes the clustering; attrs["sam_stats"] is set when SAM refinement ran.

Raises:

Type Description
ValueError

For an unsupported source, invalid parameters, a raster with too few valid pixels, or SAM refinement without sam_rgb_bands.

Source code in agribound/engines/embedding.py
def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
    """Cluster the embedding raster and polygonize the clusters.

    Parameters
    ----------
    raster_path : str
        Embedding GeoTIFF (float32, one band per embedding dimension).
    config : AgriboundConfig
        Pipeline configuration.

    Returns
    -------
    geopandas.GeoDataFrame
        Polygons with ``class_value`` (cluster label ``1..k``), in the
        raster CRS. ``attrs["engine_meta"]`` describes the clustering;
        ``attrs["sam_stats"]`` is set when SAM refinement ran.

    Raises
    ------
    ValueError
        For an unsupported source, invalid parameters, a raster with too
        few valid pixels, or SAM refinement without ``sam_rgb_bands``.
    """
    from agribound.io.raster import get_raster_info

    if config.source not in self.supported_sources:
        raise ValueError(
            f"The embedding engine needs an embedding source {self.supported_sources}, "
            f"got {config.source!r}"
        )
    params = self.resolve_params(config)
    info = get_raster_info(raster_path)
    depth = params["matryoshka_depth"]
    if depth is not None and depth > info.count:
        raise ValueError(f"matryoshka_depth={depth} but the raster has only {info.count} bands")
    if config.sam_refine:
        from agribound.engines.samgeo_engine import _rgb_band_indices

        _rgb_band_indices(config, info.count)  # fail before clustering

    cluster_path = self.cluster_cache_path(raster_path, config, params)
    meta_path = cluster_path.with_suffix(".json")

    if cluster_path.exists() and meta_path.exists():
        logger.info("Using cached embedding clusters: %s", cluster_path)
        meta = json.loads(meta_path.read_text())
        meta["cache_hit"] = True
    else:
        logger.info(
            "Clustering embeddings: %d bands, %dx%d pixels", info.count, info.width, info.height
        )
        meta = self._cluster_raster(raster_path, cluster_path, params, config)
        _atomic_write_text(meta_path, json.dumps(meta, indent=2, sort_keys=True))
        meta["cache_hit"] = False
    meta["cluster_raster"] = str(cluster_path)

    from agribound.postprocess.polygonize import polygonize_mask

    gdf = polygonize_mask(str(cluster_path), min_area_m2=config.min_field_area_m2)
    logger.info("Embedding clustering delineated %d polygons", len(gdf))

    meta["sam_refine"] = bool(config.sam_refine)
    if config.sam_refine and len(gdf) > 0:
        from agribound.engines.samgeo_engine import refine_boundaries

        gdf = refine_boundaries(gdf, raster_path, config)
        meta["sam_stats"] = gdf.attrs.get("sam_stats")
    gdf.attrs["engine_meta"] = meta
    return gdf

Ensemble

ensemble

Ensemble engine: several engines (or one engine with different models) on the same raster, combined by intersection, union or pixel vote.

Merge strategies

"intersection" (default) Successive :func:geopandas.overlay how="intersection" of the members' polygons: the output pieces are the areas every member covers, one piece per combination of overlapping polygons. Pieces are not re-merged, so where members' boundaries disagree small sliver pieces can appear next to the main piece of a field; the pipeline's area filter removes those below min_field_area_m2. "union" All members' polygons are pooled and duplicates are fused with :func:agribound.postprocess.merge.merge_polygons (IoU >= union_iou_threshold or containment >= union_containment_threshold; defaults 0.3 and 0.8). Fields that only touch stay separate. "vote" Each member's polygons are rasterised (pixel centres) onto the input raster's grid; a pixel is kept when at least min_votes members cover it, and the kept pixels are polygonised (4-connectivity). As in agribound 0.1.x, members that returned no polygons are left out of the vote (logged at WARNING and listed in vote_stats["empty_members"]), and by default min_votes = max(min(2, n), ceil(vote_threshold * n)) for the n remaining members: at least vote_threshold of them, and at least two whenever there are two or more, so that one member's false positives never pass on their own. engine_params["min_votes"] sets it directly (1..n). Members skipped after an error (on_member_error="skip") do not count either. Adjacent fields that are both kept become one polygon wherever they touch on the pixel grid, because the vote raster is binary (field / not field).

With a single member (configured, or the only one left with on_member_error="skip") that member's frame is returned as it is (all its columns), plus an engine_count column.

EnsembleEngine

Bases: DelineationEngine

Multi-engine or multi-model ensemble (see the module docstring for the strategies).

Engine parameters (config.engine_params)

engines : list[str | dict] Members; each is an engine name or a dict with "engine", optional "engine_params" and optional "label" (default: the member's engine_params["model"] or the engine name; a repeated label gets an _<index> suffix, counting up from the member's position until it is unique). Default ["delineate-anything", "ftw"]. Every member, including the defaults, must support config.source. merge_strategy : str "intersection" (default), "union" or "vote". vote_threshold : float Vote strategy: fraction in [0, 1] of the members with polygons that must agree (default 0.5); at least two members must agree whenever two or more have polygons. min_votes : int or None Vote strategy: explicit minimum number of agreeing members (1..n, where n counts the members with polygons); overrides vote_threshold. vote_resolution : float or None Vote grid cell size in CRS units. None (default) uses the input raster's own grid. union_iou_threshold, union_containment_threshold : float Union strategy duplicate criteria (defaults 0.3, 0.8). on_member_error : str "raise" (default): a failing member aborts the run. "skip": continue without it; the failure is logged as a warning and listed in engine_meta["failed_members"]. isolate_member_caches : bool Give every member its own cache directory <working dir>/ensemble/<slug> (default True), so members with different models never reuse each other's intermediates. The slug is the label with characters other than letters, digits, ., _ and - replaced by _; a slug that repeats, ignoring case, gets an _<index> suffix. Members that build their own window composites (FTW) then download them into their own directory. False shares the ensemble's cache directory.

Members receive only the engine_params of their own spec; they do not inherit the ensemble-level engine_params. Ensemble-level keys other than the ones above and the pipeline/SAM keys in :data:PIPELINE_KEYS raise ValueError (put them in a member spec). Every member runs with sam_refine=False (the pipeline refines the ensemble output instead) and is seeded with config.seed right before it runs, so its result does not depend on the member order.

Output columns: engine_count (intersection and union: number of members that ran without an error; vote: number of those with polygons) plus, per strategy, ensemble:members (intersection: all member labels; union: labels of the members whose polygons were fused, comma-separated), ensemble:n_members (union), or vote_count (maximum number of agreeing members inside the polygon; in 0.1.x this column held the constant min_votes), vote_count_mean and min_votes (vote). Member attribute columns are not carried over, except with a single member, whose frame is returned as it is plus engine_count. attrs["engine_meta"] lists every member's label, engine, parameters, polygon count, cache directory and own engine_meta.

Source code in agribound/engines/ensemble.py
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
class EnsembleEngine(DelineationEngine):
    """Multi-engine or multi-model ensemble (see the module docstring for the strategies).

    Engine parameters (``config.engine_params``)
    --------------------------------------------
    engines : list[str | dict]
        Members; each is an engine name or a dict with ``"engine"``, optional
        ``"engine_params"`` and optional ``"label"`` (default: the member's
        ``engine_params["model"]`` or the engine name; a repeated label gets
        an ``_<index>`` suffix, counting up from the member's position until
        it is unique). Default ``["delineate-anything", "ftw"]``. Every
        member, including the defaults, must support ``config.source``.
    merge_strategy : str
        ``"intersection"`` (default), ``"union"`` or ``"vote"``.
    vote_threshold : float
        Vote strategy: fraction in [0, 1] of the members with polygons that
        must agree (default 0.5); at least two members must agree whenever
        two or more have polygons.
    min_votes : int or None
        Vote strategy: explicit minimum number of agreeing members (1..n,
        where n counts the members with polygons); overrides
        *vote_threshold*.
    vote_resolution : float or None
        Vote grid cell size in CRS units. *None* (default) uses the input
        raster's own grid.
    union_iou_threshold, union_containment_threshold : float
        Union strategy duplicate criteria (defaults 0.3, 0.8).
    on_member_error : str
        ``"raise"`` (default): a failing member aborts the run. ``"skip"``:
        continue without it; the failure is logged as a warning and listed
        in ``engine_meta["failed_members"]``.
    isolate_member_caches : bool
        Give every member its own cache directory
        ``<working dir>/ensemble/<slug>`` (default *True*), so members with
        different models never reuse each other's intermediates. The slug is
        the label with characters other than letters, digits, ``.``, ``_``
        and ``-`` replaced by ``_``; a slug that repeats, ignoring case, gets
        an ``_<index>`` suffix. Members that build their own window composites
        (FTW) then download them into their own directory. *False* shares
        the ensemble's cache directory.

    Members receive only the ``engine_params`` of their own spec; they do
    not inherit the ensemble-level ``engine_params``. Ensemble-level keys
    other than the ones above and the pipeline/SAM keys in
    :data:`PIPELINE_KEYS` raise ``ValueError`` (put them in a member spec).
    Every member runs with ``sam_refine=False`` (the pipeline refines the
    ensemble output instead) and is seeded with ``config.seed`` right before
    it runs, so its result does not depend on the member order.

    Output columns: ``engine_count`` (intersection and union: number of
    members that ran without an error; vote: number of those with polygons)
    plus, per strategy,
    ``ensemble:members`` (intersection: all member labels; union: labels of
    the members whose polygons were fused, comma-separated),
    ``ensemble:n_members`` (union), or ``vote_count`` (maximum number of
    agreeing members inside the polygon; in 0.1.x this column held the
    constant ``min_votes``), ``vote_count_mean`` and ``min_votes`` (vote).
    Member attribute columns are not carried over, except with a single
    member, whose frame is returned as it is plus ``engine_count``.
    ``attrs["engine_meta"]`` lists every member's label, engine, parameters,
    polygon count, cache directory and own ``engine_meta``.
    """

    name = "ensemble"
    supported_sources = list(ENGINE_REGISTRY["ensemble"]["supported_sources"])
    requires_bands = list(ENGINE_REGISTRY["ensemble"]["requires_bands"])

    # ------------------------------------------------------------------
    # Configuration
    # ------------------------------------------------------------------

    @staticmethod
    def member_specs(config: AgriboundConfig) -> list[dict[str, Any]]:
        """Return the normalised member specs.

        Each spec is ``{"label", "engine", "engine_params", "cache_slug"}``;
        labels and cache slugs are unique.

        Raises
        ------
        ValueError
            For an empty or malformed member list, an unknown engine, a nested
            ensemble, or a member that does not support ``config.source``.
        """
        raw = (config.engine_params or {}).get("engines")
        if raw is None:
            raw = list(DEFAULT_MEMBERS)
        if not isinstance(raw, list | tuple) or not raw:
            raise ValueError("engine_params['engines'] must be a non-empty list for ensemble")
        specs: list[dict[str, Any]] = []
        used: set[str] = set()
        used_slugs: set[str] = set()
        for i, entry in enumerate(raw):
            if isinstance(entry, str):
                name, params, label = entry, {}, None
            elif isinstance(entry, dict) and isinstance(entry.get("engine"), str):
                unknown = set(entry) - {"engine", "engine_params", "label"}
                if unknown:
                    raise ValueError(f"Unknown keys {sorted(unknown)} in ensemble member {entry!r}")
                name = entry["engine"]
                params = entry.get("engine_params") or {}
                label = entry.get("label")
                if not isinstance(params, dict):
                    raise ValueError(f"engine_params of ensemble member {entry!r} must be a dict")
            else:
                raise ValueError(f"Invalid ensemble member spec: {entry!r}")
            name = name.lower().strip()
            if name not in ENGINE_REGISTRY or name == "ensemble":
                raise ValueError(
                    f"Invalid ensemble member {name!r}. Choose from "
                    f"{[n for n in ENGINE_REGISTRY if n != 'ensemble']}"
                )
            if not engine_supports_source(name, config.source):
                default = " (a default member)" if "engines" not in config.engine_params else ""
                raise ValueError(
                    f"Ensemble member {name!r}{default} does not support source "
                    f"{config.source!r} (supported: {ENGINE_REGISTRY[name]['supported_sources']}). "
                    "Set engine_params['engines'] to members that support it."
                )
            label = _unique(str(label or params.get("model") or name), used, i)
            used.add(label)
            # Distinct labels can share a slug ("a/b" and "a_b"), and macOS file
            # systems ignore case: keep the directories apart.
            slug = _unique(_slug(label, name), used_slugs, i, casefold=True)
            used_slugs.add(slug.casefold())
            specs.append(
                {
                    "label": label,
                    "engine": name,
                    "engine_params": copy.deepcopy(params),
                    "cache_slug": slug,
                }
            )
        return specs

    @staticmethod
    def resolve_params(config: AgriboundConfig) -> dict[str, Any]:
        """Return the validated ensemble-level parameters.

        Raises
        ------
        ValueError
            For unknown ensemble-level keys or invalid values.
        """
        ep = dict(config.engine_params or {})
        unknown = sorted(set(ep) - ENSEMBLE_KEYS - PIPELINE_KEYS)
        if unknown:
            raise ValueError(
                f"engine_params {unknown} are not used by the ensemble and are not passed to its "
                "members; put member settings in each member's 'engine_params' "
                "(engine_params={'engines': [{'engine': ..., 'engine_params': {...}}]})."
            )
        strategy = str(ep.get("merge_strategy", "intersection")).lower().strip()
        if strategy not in MERGE_STRATEGIES:
            raise ValueError(f"Unknown merge strategy {strategy!r}. Choose from {MERGE_STRATEGIES}")
        threshold = float(ep.get("vote_threshold", 0.5))
        if not 0.0 <= threshold <= 1.0:
            raise ValueError(f"vote_threshold must be in [0, 1], got {threshold}")
        on_error = str(ep.get("on_member_error", "raise")).lower().strip()
        if on_error not in ("raise", "skip"):
            raise ValueError(f"on_member_error must be 'raise' or 'skip', got {on_error!r}")
        resolution = ep.get("vote_resolution")
        if resolution is not None and float(resolution) <= 0:
            raise ValueError(f"vote_resolution must be > 0, got {resolution}")
        min_votes = ep.get("min_votes")
        return {
            "merge_strategy": strategy,
            "vote_threshold": threshold,
            "min_votes": None if min_votes is None else int(min_votes),
            "vote_resolution": None if resolution is None else float(resolution),
            "union_iou_threshold": float(ep.get("union_iou_threshold", 0.3)),
            "union_containment_threshold": float(ep.get("union_containment_threshold", 0.8)),
            "on_member_error": on_error,
            "isolate_member_caches": bool(ep.get("isolate_member_caches", True)),
        }

    @staticmethod
    def member_config(
        config: AgriboundConfig, spec: dict[str, Any], isolate_cache: bool = True
    ) -> AgriboundConfig:
        """Return the validated configuration a member runs with.

        Same fields as *config* except ``engine``, ``engine_params`` (the
        member's own), ``sam_refine=False`` and, with *isolate_cache*,
        ``cache_dir=<working dir>/ensemble/<spec["cache_slug"]>`` (the slug
        of the label when the spec has no ``cache_slug``).
        """
        overrides: dict[str, Any] = {
            "engine": spec["engine"],
            "engine_params": copy.deepcopy(spec["engine_params"]),
            "sam_refine": False,
        }
        if isolate_cache:
            slug = spec.get("cache_slug") or _slug(spec["label"], spec["engine"])
            overrides["cache_dir"] = str(Path(config.get_working_dir()) / "ensemble" / slug)
        return config.merged(**overrides)

    # ------------------------------------------------------------------
    # Engine API
    # ------------------------------------------------------------------

    @classmethod
    def prefetch(cls, config: AgriboundConfig) -> list[str]:
        """Prefetch every member's weights (duplicates removed, order kept)."""
        paths: list[str] = []
        for spec in cls.member_specs(config):
            member = cls.member_config(config, spec, isolate_cache=False)
            for path in get_engine_class(spec["engine"]).prefetch(member) or []:
                if str(path) not in paths:
                    paths.append(str(path))
        return paths

    @classmethod
    def stage_inputs(cls, config: AgriboundConfig, raster_path: str) -> dict[str, Any]:
        """Build the inputs that members build themselves, without running them.

        For every member whose engine class has a ``stage_inputs`` method
        (currently FTW: :meth:`agribound.engines.ftw.FTWEngine.stage_inputs`,
        the two seasonal window composites), calls it with the member's own
        configuration (:meth:`member_config`, including its isolated cache
        directory) and *raster_path*, so a later :meth:`delineate` with the
        same configuration finds the inputs in the member caches without
        network access. Used by :mod:`agribound.hpc.tiles` to stage ensemble
        tiles.

        Parameters
        ----------
        config : AgriboundConfig
            Ensemble configuration.
        raster_path : str
            Input raster shared by the members.

        Returns
        -------
        dict
            ``members`` (label -> the member's ``stage_inputs`` result),
            ``rasters`` (all staged rasters, member order) and
            ``failed_members`` (``{"label", "engine", "error"}``; only with
            ``on_member_error="skip"``, as in :meth:`delineate`).

        Raises
        ------
        Exception
            A member's staging error when ``on_member_error="raise"``.
        """
        specs = cls.member_specs(config)
        params = cls.resolve_params(config)
        members: dict[str, Any] = {}
        rasters: list[str] = []
        failed: list[dict[str, str]] = []
        for spec in specs:
            stage = getattr(get_engine_class(spec["engine"]), "stage_inputs", None)
            if stage is None:
                continue
            member = cls.member_config(config, spec, params["isolate_member_caches"])
            try:
                staged = stage(member, raster_path)
            except Exception as exc:
                if params["on_member_error"] == "raise":
                    exc.add_note(f"(ensemble member {spec['label']!r}, engine {spec['engine']!r})")
                    raise
                logger.warning(
                    "Staging the inputs of ensemble member %s failed; it will be retried (and "
                    "skipped if it fails) during delineation: %s",
                    spec["label"],
                    exc,
                )
                failed.append(
                    {
                        "label": spec["label"],
                        "engine": spec["engine"],
                        "error": f"{type(exc).__name__}: {exc}",
                    }
                )
                continue
            members[spec["label"]] = staged
            rasters.extend(str(p) for p in staged.get("rasters", []))
        return {"members": members, "rasters": rasters, "failed_members": failed}

    def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
        """Run every member on *raster_path* and merge the results.

        Parameters
        ----------
        raster_path : str
            Input GeoTIFF shared by all members.
        config : AgriboundConfig
            Pipeline configuration (``engine="ensemble"``).

        Returns
        -------
        geopandas.GeoDataFrame
            Merged polygons in the raster CRS, with ``attrs["engine_meta"]``.

        Raises
        ------
        ValueError
            For invalid ensemble parameters or members.
        RuntimeError
            If every member failed (``on_member_error="skip"``).
        """
        from agribound._repro import seed_everything
        from agribound.io.raster import get_raster_info

        specs = self.member_specs(config)
        params = self.resolve_params(config)
        info = get_raster_info(raster_path)
        target_crs = info.crs

        results: dict[str, gpd.GeoDataFrame] = {}
        members_meta: list[dict[str, Any]] = []
        failed: list[dict[str, str]] = []
        for i, spec in enumerate(specs):
            label = spec["label"]
            member = self.member_config(config, spec, params["isolate_member_caches"])
            logger.info(
                "Ensemble [%d/%d]: running %s (%s)", i + 1, len(specs), label, spec["engine"]
            )
            seed_everything(member.seed)
            try:
                gdf = get_engine(spec["engine"]).delineate(raster_path, member)
            except Exception as exc:
                if params["on_member_error"] == "raise":
                    exc.add_note(f"(ensemble member {label!r}, engine {spec['engine']!r})")
                    raise
                logger.warning("Ensemble member %s failed and is skipped: %s", label, exc)
                failed.append(
                    {
                        "label": label,
                        "engine": spec["engine"],
                        "error": f"{type(exc).__name__}: {exc}",
                    }
                )
                continue
            gdf = _as_frame(gdf, target_crs)
            results[label] = gdf
            members_meta.append(
                {
                    "label": label,
                    "engine": spec["engine"],
                    "engine_params": spec["engine_params"],
                    "n_polygons": int(len(gdf)),
                    "cache_dir": str(member.get_working_dir()),
                    "seed": int(member.seed),
                    "engine_meta": copy.deepcopy(gdf.attrs.get("engine_meta")),
                }
            )
            logger.info("%s produced %d polygons", label, len(gdf))

        if not results:
            raise RuntimeError(
                "All ensemble members failed:\n"
                + "\n".join(f"  - {f['label']}: {f['error']}" for f in failed)
            )

        strategy = params["merge_strategy"]
        meta: dict[str, Any] = {
            "backend": "ensemble",
            "merge_strategy": strategy,
            "n_members": len(results),
            "members": members_meta,
            "failed_members": failed,
            "isolate_member_caches": params["isolate_member_caches"],
        }
        if len(results) == 1:
            merged = next(iter(results.values())).copy()
            merged["engine_count"] = 1
            meta["note"] = "single member: its polygons are returned unchanged"
        elif strategy == "union":
            merged = self._merge_union(
                results,
                iou_threshold=params["union_iou_threshold"],
                containment_threshold=params["union_containment_threshold"],
            )
            meta["union_iou_threshold"] = params["union_iou_threshold"]
            meta["union_containment_threshold"] = params["union_containment_threshold"]
        elif strategy == "intersection":
            merged = self._merge_intersection(results)
        else:
            grid = None
            if params["vote_resolution"] is None:
                grid = (info.crs, info.transform, info.width, info.height)
            merged = self._merge_vote(
                results,
                params["vote_threshold"],
                min_votes=params["min_votes"],
                resolution=params["vote_resolution"],
                grid=grid,
            )
            meta["vote_threshold"] = params["vote_threshold"]
            meta["vote"] = copy.deepcopy(merged.attrs.get("vote_stats"))
        merged.attrs = {"engine_meta": meta}
        return merged

    # ------------------------------------------------------------------
    # Merge strategies (static; also usable on saved member outputs)
    # ------------------------------------------------------------------

    @staticmethod
    def _merge_union(
        results: dict[str, gpd.GeoDataFrame],
        iou_threshold: float = 0.3,
        containment_threshold: float = 0.8,
    ) -> gpd.GeoDataFrame:
        """Pool all members' polygons and fuse duplicates (see module docstring)."""
        from agribound.postprocess.merge import merge_polygons

        frames, crs = _aligned(results)
        labels = list(frames)
        rows = []
        for label in labels:
            gdf = frames[label]
            geoms = [g for g in gdf.geometry if g is not None and not g.is_empty]
            rows.append(
                gpd.GeoDataFrame({"_member": [label] * len(geoms)}, geometry=geoms, crs=crs)
            )
        pooled = gpd.GeoDataFrame(pd.concat(rows, ignore_index=True), geometry="geometry", crs=crs)
        if len(pooled) == 0:
            return _empty(crs, [MEMBERS_COLUMN, N_MEMBERS_COLUMN, "engine_count"])
        merged, groups = merge_polygons(
            pooled,
            iou_threshold=iou_threshold,
            containment_threshold=containment_threshold,
            return_groups=True,
        )
        member_of = pooled["_member"].to_numpy()
        names = [sorted({str(member_of[p]) for p in group}) for group in groups]
        out = gpd.GeoDataFrame(
            {
                MEMBERS_COLUMN: [",".join(n) for n in names],
                N_MEMBERS_COLUMN: [len(n) for n in names],
                "engine_count": len(results),
            },
            geometry=list(merged.geometry),
            crs=crs,
        )
        logger.info("Union merge: %d polygons from %d pooled", len(out), len(pooled))
        return out

    @staticmethod
    def _merge_intersection(results: dict[str, gpd.GeoDataFrame]) -> gpd.GeoDataFrame:
        """Areas covered by every member, via successive overlays (see module docstring)."""
        frames, crs = _aligned(results)
        labels = list(frames)
        columns = [MEMBERS_COLUMN, "engine_count"]
        base = _geometry_only(frames[labels[0]])
        for label in labels[1:]:
            other = _geometry_only(frames[label])
            if len(base) == 0 or len(other) == 0:
                base = base.iloc[0:0]
                break
            base = gpd.overlay(base, other, how="intersection", keep_geom_type=True)
            base = _geometry_only(base)
        base = base[~base.geometry.isna() & ~base.geometry.is_empty]
        if len(base) == 0:
            logger.warning("Intersection merge produced no overlapping polygons")
            return _empty(crs, columns)
        n_out = len(base)
        out = gpd.GeoDataFrame(
            {
                MEMBERS_COLUMN: [",".join(sorted(labels))] * n_out,
                "engine_count": [len(results)] * n_out,
            },
            geometry=list(base.geometry),
            crs=crs,
        )
        logger.info("Intersection merge: %d polygons", len(out))
        return out

    @staticmethod
    def _merge_vote(
        results: dict[str, gpd.GeoDataFrame],
        threshold: float = 0.5,
        *,
        min_votes: int | None = None,
        resolution: float | None = None,
        grid: tuple[Any, Any, int, int] | None = None,
    ) -> gpd.GeoDataFrame:
        """Pixel vote (see module docstring).

        Parameters
        ----------
        results : dict[str, geopandas.GeoDataFrame]
            Member polygons by label. Members without polygons are left out
            of the vote (``n`` counts the others), as in agribound 0.1.x.
        threshold : float
            Vote threshold in [0, 1] (ignored when *min_votes* is given):
            ``min_votes = max(min(2, n), ceil(threshold * n))``.
        min_votes : int or None
            Explicit minimum number of agreeing members (1..n).
        resolution : float or None
            Cell size in CRS units for a grid over the members' combined
            extent. Used when *grid* is *None*; *None* then means 10 units in
            a projected CRS or 1e-4 degrees in a geographic one (the 0.1.x grid).
        grid : tuple or None
            ``(crs, transform, width, height)`` of the vote raster, e.g. the
            input raster's grid (what :meth:`delineate` passes).

        Returns
        -------
        geopandas.GeoDataFrame
            Polygons with ``vote_count`` (max agreement inside), ``vote_count_mean``,
            ``min_votes`` and ``engine_count`` (``n``); ``attrs["vote_stats"]``
            records ``min_votes``, the rule, ``threshold``, ``n_members``
            (``n``), ``n_members_total``, ``empty_members`` and the grid
            (*None* when no member has polygons). An explicit *min_votes*
            larger than ``n`` gives an empty result (logged at WARNING).

        Raises
        ------
        ValueError
            For no member results, a threshold outside [0, 1] or *min_votes*
            outside 1..n.
        """
        import rasterio
        from rasterio.features import rasterize, shapes
        from scipy import ndimage
        from shapely.geometry import shape as shapely_shape

        from agribound.postprocess.simplify import make_polygonal

        n_total = len(results)
        if n_total == 0:
            raise ValueError("_merge_vote needs at least one member result")
        if min_votes is None and not 0.0 <= threshold <= 1.0:
            raise ValueError(f"vote threshold must be in [0, 1], got {threshold}")
        if min_votes is not None and not 1 <= int(min_votes) <= n_total:
            raise ValueError(f"min_votes must be in [1, {n_total}], got {min_votes}")
        columns = ["vote_count", "vote_count_mean", "min_votes", "engine_count"]

        # As in 0.1.x, members without polygons are left out of the vote.
        empty_members = [k for k, v in results.items() if not _has_polygons(v)]
        voting = {k: v for k, v in results.items() if k not in empty_members}
        n = len(voting)
        if min_votes is None:
            rule = "max(min(2, n), ceil(threshold * n)) (agribound 0.1.x)"
            min_votes = _min_votes(n, threshold) if n else None
        else:
            rule = "explicit min_votes"
            min_votes = int(min_votes)
        stats: dict[str, Any] = {
            "min_votes": min_votes,
            "rule": rule,
            "n_members": n,
            "n_members_total": n_total,
            "empty_members": empty_members,
            "threshold": threshold,
            "grid": None,
        }
        if empty_members:
            logger.warning(
                "Vote merge: %d of %d members returned no polygons and are left out of the "
                "vote (%s)",
                len(empty_members),
                n_total,
                ", ".join(empty_members),
            )
        if n == 0:
            crs = grid[0] if grid is not None else _aligned(results)[1]
            out = _empty(crs, columns)
            out.attrs["vote_stats"] = stats
            return out
        if min_votes > n:
            logger.warning(
                "Vote merge: min_votes=%d but only %d member(s) returned polygons; the "
                "result is empty",
                min_votes,
                n,
            )

        if grid is not None:
            crs, transform, width, height = grid
            frames = {
                k: (v.to_crs(crs) if v.crs is not None and v.crs != crs else v)
                for k, v in voting.items()
            }
        else:
            frames, crs = _aligned(voting)
            bounds = np.array([f.total_bounds for f in frames.values()])
            minx, miny = np.nanmin(bounds[:, 0]), np.nanmin(bounds[:, 1])
            maxx, maxy = np.nanmax(bounds[:, 2]), np.nanmax(bounds[:, 3])
            if resolution is None:
                resolution = 1e-4 if crs is not None and crs.is_geographic else 10.0
            width = int(np.ceil((maxx - minx) / resolution))
            height = int(np.ceil((maxy - miny) / resolution))
            if width == 0 or height == 0:
                out = _empty(crs, columns)
                out.attrs["vote_stats"] = stats
                return out
            transform = rasterio.transform.from_origin(minx, maxy, resolution, resolution)
        stats["grid"] = {
            "crs": str(crs) if crs is not None else None,
            "transform": list(transform)[:6],
            "width": int(width),
            "height": int(height),
        }

        votes = np.zeros((height, width), dtype=np.uint16)
        for gdf in frames.values():
            geoms = [make_polygonal(g) for g in gdf.geometry if g is not None and not g.is_empty]
            burn = [(g, 1) for g in geoms if g is not None and not g.is_empty]
            if burn:
                votes += rasterize(
                    burn, out_shape=(height, width), transform=transform, fill=0, dtype=np.uint8
                )

        consensus = votes >= min_votes
        labels, n_labels = ndimage.label(consensus, structure=[[0, 1, 0], [1, 1, 1], [0, 1, 0]])
        if n_labels == 0:
            out = _empty(crs, columns)
            out.attrs["vote_stats"] = stats
            return out
        index = np.arange(1, n_labels + 1)
        max_votes = np.asarray(ndimage.maximum(votes, labels, index), dtype=int)
        mean_votes = np.asarray(ndimage.mean(votes, labels, index), dtype=float)

        polys: dict[int, Any] = {}
        for geom, val in shapes(
            labels.astype(np.int32), mask=consensus, transform=transform, connectivity=4
        ):
            polys[int(val)] = shapely_shape(geom)
        ids = sorted(polys)
        out = gpd.GeoDataFrame(
            {
                "vote_count": [int(max_votes[i - 1]) for i in ids],
                "vote_count_mean": [float(mean_votes[i - 1]) for i in ids],
                "min_votes": min_votes,
                "engine_count": n,
            },
            geometry=[polys[i] for i in ids],
            crs=crs,
        )
        out.attrs["vote_stats"] = stats
        logger.info("Vote merge: %d polygons (at least %d of %d members)", len(out), min_votes, n)
        return out
member_specs staticmethod
member_specs(config: AgriboundConfig) -> list[dict[str, Any]]

Return the normalised member specs.

Each spec is {"label", "engine", "engine_params", "cache_slug"}; labels and cache slugs are unique.

Raises:

Type Description
ValueError

For an empty or malformed member list, an unknown engine, a nested ensemble, or a member that does not support config.source.

Source code in agribound/engines/ensemble.py
@staticmethod
def member_specs(config: AgriboundConfig) -> list[dict[str, Any]]:
    """Return the normalised member specs.

    Each spec is ``{"label", "engine", "engine_params", "cache_slug"}``;
    labels and cache slugs are unique.

    Raises
    ------
    ValueError
        For an empty or malformed member list, an unknown engine, a nested
        ensemble, or a member that does not support ``config.source``.
    """
    raw = (config.engine_params or {}).get("engines")
    if raw is None:
        raw = list(DEFAULT_MEMBERS)
    if not isinstance(raw, list | tuple) or not raw:
        raise ValueError("engine_params['engines'] must be a non-empty list for ensemble")
    specs: list[dict[str, Any]] = []
    used: set[str] = set()
    used_slugs: set[str] = set()
    for i, entry in enumerate(raw):
        if isinstance(entry, str):
            name, params, label = entry, {}, None
        elif isinstance(entry, dict) and isinstance(entry.get("engine"), str):
            unknown = set(entry) - {"engine", "engine_params", "label"}
            if unknown:
                raise ValueError(f"Unknown keys {sorted(unknown)} in ensemble member {entry!r}")
            name = entry["engine"]
            params = entry.get("engine_params") or {}
            label = entry.get("label")
            if not isinstance(params, dict):
                raise ValueError(f"engine_params of ensemble member {entry!r} must be a dict")
        else:
            raise ValueError(f"Invalid ensemble member spec: {entry!r}")
        name = name.lower().strip()
        if name not in ENGINE_REGISTRY or name == "ensemble":
            raise ValueError(
                f"Invalid ensemble member {name!r}. Choose from "
                f"{[n for n in ENGINE_REGISTRY if n != 'ensemble']}"
            )
        if not engine_supports_source(name, config.source):
            default = " (a default member)" if "engines" not in config.engine_params else ""
            raise ValueError(
                f"Ensemble member {name!r}{default} does not support source "
                f"{config.source!r} (supported: {ENGINE_REGISTRY[name]['supported_sources']}). "
                "Set engine_params['engines'] to members that support it."
            )
        label = _unique(str(label or params.get("model") or name), used, i)
        used.add(label)
        # Distinct labels can share a slug ("a/b" and "a_b"), and macOS file
        # systems ignore case: keep the directories apart.
        slug = _unique(_slug(label, name), used_slugs, i, casefold=True)
        used_slugs.add(slug.casefold())
        specs.append(
            {
                "label": label,
                "engine": name,
                "engine_params": copy.deepcopy(params),
                "cache_slug": slug,
            }
        )
    return specs
resolve_params staticmethod
resolve_params(config: AgriboundConfig) -> dict[str, Any]

Return the validated ensemble-level parameters.

Raises:

Type Description
ValueError

For unknown ensemble-level keys or invalid values.

Source code in agribound/engines/ensemble.py
@staticmethod
def resolve_params(config: AgriboundConfig) -> dict[str, Any]:
    """Return the validated ensemble-level parameters.

    Raises
    ------
    ValueError
        For unknown ensemble-level keys or invalid values.
    """
    ep = dict(config.engine_params or {})
    unknown = sorted(set(ep) - ENSEMBLE_KEYS - PIPELINE_KEYS)
    if unknown:
        raise ValueError(
            f"engine_params {unknown} are not used by the ensemble and are not passed to its "
            "members; put member settings in each member's 'engine_params' "
            "(engine_params={'engines': [{'engine': ..., 'engine_params': {...}}]})."
        )
    strategy = str(ep.get("merge_strategy", "intersection")).lower().strip()
    if strategy not in MERGE_STRATEGIES:
        raise ValueError(f"Unknown merge strategy {strategy!r}. Choose from {MERGE_STRATEGIES}")
    threshold = float(ep.get("vote_threshold", 0.5))
    if not 0.0 <= threshold <= 1.0:
        raise ValueError(f"vote_threshold must be in [0, 1], got {threshold}")
    on_error = str(ep.get("on_member_error", "raise")).lower().strip()
    if on_error not in ("raise", "skip"):
        raise ValueError(f"on_member_error must be 'raise' or 'skip', got {on_error!r}")
    resolution = ep.get("vote_resolution")
    if resolution is not None and float(resolution) <= 0:
        raise ValueError(f"vote_resolution must be > 0, got {resolution}")
    min_votes = ep.get("min_votes")
    return {
        "merge_strategy": strategy,
        "vote_threshold": threshold,
        "min_votes": None if min_votes is None else int(min_votes),
        "vote_resolution": None if resolution is None else float(resolution),
        "union_iou_threshold": float(ep.get("union_iou_threshold", 0.3)),
        "union_containment_threshold": float(ep.get("union_containment_threshold", 0.8)),
        "on_member_error": on_error,
        "isolate_member_caches": bool(ep.get("isolate_member_caches", True)),
    }
member_config staticmethod
member_config(config: AgriboundConfig, spec: dict[str, Any], isolate_cache: bool = True) -> AgriboundConfig

Return the validated configuration a member runs with.

Same fields as config except engine, engine_params (the member's own), sam_refine=False and, with isolate_cache, cache_dir=<working dir>/ensemble/<spec["cache_slug"]> (the slug of the label when the spec has no cache_slug).

Source code in agribound/engines/ensemble.py
@staticmethod
def member_config(
    config: AgriboundConfig, spec: dict[str, Any], isolate_cache: bool = True
) -> AgriboundConfig:
    """Return the validated configuration a member runs with.

    Same fields as *config* except ``engine``, ``engine_params`` (the
    member's own), ``sam_refine=False`` and, with *isolate_cache*,
    ``cache_dir=<working dir>/ensemble/<spec["cache_slug"]>`` (the slug
    of the label when the spec has no ``cache_slug``).
    """
    overrides: dict[str, Any] = {
        "engine": spec["engine"],
        "engine_params": copy.deepcopy(spec["engine_params"]),
        "sam_refine": False,
    }
    if isolate_cache:
        slug = spec.get("cache_slug") or _slug(spec["label"], spec["engine"])
        overrides["cache_dir"] = str(Path(config.get_working_dir()) / "ensemble" / slug)
    return config.merged(**overrides)
prefetch classmethod
prefetch(config: AgriboundConfig) -> list[str]

Prefetch every member's weights (duplicates removed, order kept).

Source code in agribound/engines/ensemble.py
@classmethod
def prefetch(cls, config: AgriboundConfig) -> list[str]:
    """Prefetch every member's weights (duplicates removed, order kept)."""
    paths: list[str] = []
    for spec in cls.member_specs(config):
        member = cls.member_config(config, spec, isolate_cache=False)
        for path in get_engine_class(spec["engine"]).prefetch(member) or []:
            if str(path) not in paths:
                paths.append(str(path))
    return paths
stage_inputs classmethod
stage_inputs(config: AgriboundConfig, raster_path: str) -> dict[str, Any]

Build the inputs that members build themselves, without running them.

For every member whose engine class has a stage_inputs method (currently FTW: :meth:agribound.engines.ftw.FTWEngine.stage_inputs, the two seasonal window composites), calls it with the member's own configuration (:meth:member_config, including its isolated cache directory) and raster_path, so a later :meth:delineate with the same configuration finds the inputs in the member caches without network access. Used by :mod:agribound.hpc.tiles to stage ensemble tiles.

Parameters:

Name Type Description Default
config AgriboundConfig

Ensemble configuration.

required
raster_path str

Input raster shared by the members.

required

Returns:

Type Description
dict

members (label -> the member's stage_inputs result), rasters (all staged rasters, member order) and failed_members ({"label", "engine", "error"}; only with on_member_error="skip", as in :meth:delineate).

Raises:

Type Description
Exception

A member's staging error when on_member_error="raise".

Source code in agribound/engines/ensemble.py
@classmethod
def stage_inputs(cls, config: AgriboundConfig, raster_path: str) -> dict[str, Any]:
    """Build the inputs that members build themselves, without running them.

    For every member whose engine class has a ``stage_inputs`` method
    (currently FTW: :meth:`agribound.engines.ftw.FTWEngine.stage_inputs`,
    the two seasonal window composites), calls it with the member's own
    configuration (:meth:`member_config`, including its isolated cache
    directory) and *raster_path*, so a later :meth:`delineate` with the
    same configuration finds the inputs in the member caches without
    network access. Used by :mod:`agribound.hpc.tiles` to stage ensemble
    tiles.

    Parameters
    ----------
    config : AgriboundConfig
        Ensemble configuration.
    raster_path : str
        Input raster shared by the members.

    Returns
    -------
    dict
        ``members`` (label -> the member's ``stage_inputs`` result),
        ``rasters`` (all staged rasters, member order) and
        ``failed_members`` (``{"label", "engine", "error"}``; only with
        ``on_member_error="skip"``, as in :meth:`delineate`).

    Raises
    ------
    Exception
        A member's staging error when ``on_member_error="raise"``.
    """
    specs = cls.member_specs(config)
    params = cls.resolve_params(config)
    members: dict[str, Any] = {}
    rasters: list[str] = []
    failed: list[dict[str, str]] = []
    for spec in specs:
        stage = getattr(get_engine_class(spec["engine"]), "stage_inputs", None)
        if stage is None:
            continue
        member = cls.member_config(config, spec, params["isolate_member_caches"])
        try:
            staged = stage(member, raster_path)
        except Exception as exc:
            if params["on_member_error"] == "raise":
                exc.add_note(f"(ensemble member {spec['label']!r}, engine {spec['engine']!r})")
                raise
            logger.warning(
                "Staging the inputs of ensemble member %s failed; it will be retried (and "
                "skipped if it fails) during delineation: %s",
                spec["label"],
                exc,
            )
            failed.append(
                {
                    "label": spec["label"],
                    "engine": spec["engine"],
                    "error": f"{type(exc).__name__}: {exc}",
                }
            )
            continue
        members[spec["label"]] = staged
        rasters.extend(str(p) for p in staged.get("rasters", []))
    return {"members": members, "rasters": rasters, "failed_members": failed}
delineate
delineate(raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame

Run every member on raster_path and merge the results.

Parameters:

Name Type Description Default
raster_path str

Input GeoTIFF shared by all members.

required
config AgriboundConfig

Pipeline configuration (engine="ensemble").

required

Returns:

Type Description
GeoDataFrame

Merged polygons in the raster CRS, with attrs["engine_meta"].

Raises:

Type Description
ValueError

For invalid ensemble parameters or members.

RuntimeError

If every member failed (on_member_error="skip").

Source code in agribound/engines/ensemble.py
def delineate(self, raster_path: str, config: AgriboundConfig) -> gpd.GeoDataFrame:
    """Run every member on *raster_path* and merge the results.

    Parameters
    ----------
    raster_path : str
        Input GeoTIFF shared by all members.
    config : AgriboundConfig
        Pipeline configuration (``engine="ensemble"``).

    Returns
    -------
    geopandas.GeoDataFrame
        Merged polygons in the raster CRS, with ``attrs["engine_meta"]``.

    Raises
    ------
    ValueError
        For invalid ensemble parameters or members.
    RuntimeError
        If every member failed (``on_member_error="skip"``).
    """
    from agribound._repro import seed_everything
    from agribound.io.raster import get_raster_info

    specs = self.member_specs(config)
    params = self.resolve_params(config)
    info = get_raster_info(raster_path)
    target_crs = info.crs

    results: dict[str, gpd.GeoDataFrame] = {}
    members_meta: list[dict[str, Any]] = []
    failed: list[dict[str, str]] = []
    for i, spec in enumerate(specs):
        label = spec["label"]
        member = self.member_config(config, spec, params["isolate_member_caches"])
        logger.info(
            "Ensemble [%d/%d]: running %s (%s)", i + 1, len(specs), label, spec["engine"]
        )
        seed_everything(member.seed)
        try:
            gdf = get_engine(spec["engine"]).delineate(raster_path, member)
        except Exception as exc:
            if params["on_member_error"] == "raise":
                exc.add_note(f"(ensemble member {label!r}, engine {spec['engine']!r})")
                raise
            logger.warning("Ensemble member %s failed and is skipped: %s", label, exc)
            failed.append(
                {
                    "label": label,
                    "engine": spec["engine"],
                    "error": f"{type(exc).__name__}: {exc}",
                }
            )
            continue
        gdf = _as_frame(gdf, target_crs)
        results[label] = gdf
        members_meta.append(
            {
                "label": label,
                "engine": spec["engine"],
                "engine_params": spec["engine_params"],
                "n_polygons": int(len(gdf)),
                "cache_dir": str(member.get_working_dir()),
                "seed": int(member.seed),
                "engine_meta": copy.deepcopy(gdf.attrs.get("engine_meta")),
            }
        )
        logger.info("%s produced %d polygons", label, len(gdf))

    if not results:
        raise RuntimeError(
            "All ensemble members failed:\n"
            + "\n".join(f"  - {f['label']}: {f['error']}" for f in failed)
        )

    strategy = params["merge_strategy"]
    meta: dict[str, Any] = {
        "backend": "ensemble",
        "merge_strategy": strategy,
        "n_members": len(results),
        "members": members_meta,
        "failed_members": failed,
        "isolate_member_caches": params["isolate_member_caches"],
    }
    if len(results) == 1:
        merged = next(iter(results.values())).copy()
        merged["engine_count"] = 1
        meta["note"] = "single member: its polygons are returned unchanged"
    elif strategy == "union":
        merged = self._merge_union(
            results,
            iou_threshold=params["union_iou_threshold"],
            containment_threshold=params["union_containment_threshold"],
        )
        meta["union_iou_threshold"] = params["union_iou_threshold"]
        meta["union_containment_threshold"] = params["union_containment_threshold"]
    elif strategy == "intersection":
        merged = self._merge_intersection(results)
    else:
        grid = None
        if params["vote_resolution"] is None:
            grid = (info.crs, info.transform, info.width, info.height)
        merged = self._merge_vote(
            results,
            params["vote_threshold"],
            min_votes=params["min_votes"],
            resolution=params["vote_resolution"],
            grid=grid,
        )
        meta["vote_threshold"] = params["vote_threshold"]
        meta["vote"] = copy.deepcopy(merged.attrs.get("vote_stats"))
    merged.attrs = {"engine_meta": meta}
    return merged

SAM refinement

SAM 3 backends are untested

sam_backend="sam3" and "sam3-hf" have not been run end to end with agribound 1.0.1 (the facebook/sam3 weights are gated); a WARNING is logged when one is loaded. See SAM refinement.

samgeo_engine

Box-prompted SAM refinement of field boundaries.

This is a post-processing stage, not a delineation engine: every polygon's bounding box is given to a Segment Anything model as a single-instance box prompt, and the polygon is replaced by the mask SAM returns, unless that mask covers too little of it (step 7). The pipeline runs it after delineation when config.sam_refine is True (the embedding engine calls it itself).

Algorithm (:func:refine_boundaries)
  1. Polygons are reprojected to the raster CRS. The raster must not be rotated or sheared.
  2. Gating. Missing or empty geometries are skipped. So is every polygon whose bounding box extends more than half a pixel (:data:EDGE_TOLERANCE_PX) beyond the raster (counted in n_skipped_outside): SAM would see only part of such a field, and a mask of the visible part would truncate it. Polygons delineated from the raster itself normally lie inside it. Of the remaining polygons, one is refined only if both sides of its padded bounding box are at least config.sam_min_crop_px pixels, where the padded side is floor(side_px * (1 + 2 * config.sam_crop_padding)) and side_px is the bounding-box width (height) divided by the pixel width (height). :func:crop_window_px and :func:is_refinable implement exactly this test (pass raster_bounds to :func:is_refinable to include the inside-the-raster test). Skipped polygons keep their geometry. Boxes within the half-pixel tolerance, and padded boxes, are clipped to the raster.
  3. Image. The canonical R, G, B bands of config.source (with config.bands taking precedence, or engine_params["sam_rgb_bands"]) are converted to uint8 with one scene-wide percentile stretch: the 1st and 99th percentiles of each band are computed by :func:agribound.io.raster.percentile_stretch_uint8 (valid, finite, positive pixels) on a nearest-neighbour decimated read of at most 4096 pixels per side, and every window is stretched with those same bounds. uint8 rasters are used as-is. For embedding rasters (signed values) the percentiles are taken over all finite values.
  4. Windows. Polygons whose padded box (clipped to the raster) fits in window_px - window_px // 2 pixels on both axes are assigned to a grid of window_px x window_px windows with a stride of window_px // 2 (windows at the raster edge are shifted inwards to keep their full size); each window is encoded once and all its boxes are decoded in batches of batch_size. Larger polygons get their own square window, with the side of their padded box's longer axis, centred on the box. A field window whose side exceeds max_window_px = max(2 * window_px, 2048) is read with nearest-neighbour decimation so that its longer side is max_window_px, which bounds memory (SAM downsamples it to its input size in any case). SAM resizes every window to a fixed square input without keeping the aspect ratio (1024 x 1024 px for SAM 2/2.1, 1008 x 1008 px for SAM 3), so windows that are not square (the raster is smaller than the window on one axis, or a field window is clipped by the raster) are padded with black (0) pixels on the right and bottom to a square first; box coordinates are unaffected. A full grid window is encoded at about its native scale. This differs from agribound 0.1.x, which encoded each field's padded crop on its own, so that SAM upsampled a 64 px crop 16-fold; a smaller engine_params["sam_window_px"] (at least 2 * sam_min_crop_px) restores part of that zoom at the cost of more encoder passes.
  5. Masks. One mask per box (multimask_output=False). As in 0.1.x, where the image was the field's padded crop, only the part of the mask inside the field's padded box (rounded outwards to whole pixels and clipped to the raster) is used. It is vectorised with :func:rasterio.features.shapes, the largest polygon is kept (holes included), repaired if invalid, and reprojected to the input CRS. Boxes are passed as continuous pixel coordinates (not truncated). Steps 6 and 7 decide whether the mask replaces the input polygon.
  6. Overlaps (engine_params["sam_overlaps"], default "trim"). A mask may grow over a neighbouring polygon, and both would be kept. With "trim" a refined polygon never takes area that another input polygon covered and it did not (:func:trim_refinement_overlaps), and where two refined masks grew over the same new area, the one with the higher SAM score keeps it. The overlap between the refined polygons and the others is therefore never larger than between the input polygons, so SAM adds no overlap to an engine output without overlaps (Delineate-Anything resolves them). The trimmed mask keeps its largest part (step 7 then tests it); a mask with nothing left keeps the input geometry and counts as failed. "keep" keeps the masks as SAM drew them (with sam_min_coverage=0, the behaviour before 1.0.0), so outputs may overlap. Trade-off, measured on 2026-09-28 with Delineate-Anything (large_v2) + SAM 2 (sam2-hiera-large, MPS) on the Namoi test area (Sentinel-2, 2023; 16 of 230 polygons refined): the post-processed output had 3.59 ha of overlap with "keep" (3.45 ha between a refined polygon and a neighbour) and 0.20 ha with "trim", against 0.18 ha without SAM. With "keep" one refined mask grew over the neighbours of one of the four reference fields in the area, raising that field's best IoU from 0.636 (no SAM) to 0.773; with "trim" it is 0.676 (the other three fields: unchanged or within 0.02). SAM's growth over an engine boundary can be a correction (a field split in two) or a leak into a real neighbour; the default keeps the engine's boundaries between polygons.
  7. Coverage (engine_params["sam_min_coverage"], default :data:DEFAULT_MIN_COVERAGE = 0.5). SAM returns one object per box, so when an input polygon holds several fields (an embedding cluster, or fields an engine merged) the mask can follow one of them, and the rest of the polygon's area would be left without a polygon. A mask that covers less than sam_min_coverage of its input polygon, area(mask & input) / area(input) after step 6 (an invalid input is repaired first), is not used: the polygon keeps its input geometry, agribound:sam_refined is False and agribound:sam_score NaN, and n_low_coverage counts it. With "trim" each mask is tested right after its trim, in step 6's score order, so a rejected mask claims no area from the lower-scoring masks. An input without area is never rejected; 0 turns the test off (the 1.0.0 behaviour). Measured on 2026-09-29 (SAM 2 sam2-hiera-large, MPS, outputs after the area filter, smoothing and simplification), no test -> 0.5: example 15's Sentinel-2 refinement of the TESSERA (Google) crop polygons left 8 -> 1 (14 -> 2) of 29 checked centre pivots less than half covered and lost 15.7 -> 7.4 % (24.1 -> 8.1 %) of the input area (48 of 510, 95 of 439 masks rejected); example 13's input (Delineate-Anything, Sentinel-2, 67 masks, each covering >= 92 % of its polygon) did not change. On example 14's DINOv3 NAIP (SPOT) polygons of Lea County the reference fields less than half covered went from 35 (67) without SAM to 80 (105) with SAM and 53 (85) with 0.5, and F1 from 0.604 (0.423) to 0.609 (0.479) and 0.590 (0.445): a rejected mask often matched one of the several fields its polygon held. 0.7 restored more coverage (no pivot missing; 40 (77) fields) at an F1 of 0.595 (0.418).
Backends (config.sam_backend)
  • "sam2": samgeo.SamGeo2(model_id, automatic=False).predictor (SAM2ImagePredictor); SAM 2.0 checkpoints facebook/sam2-hiera-*.
  • "sam2.1": sam2.sam2_image_predictor.SAM2ImagePredictor.from_pretrained with facebook/sam2.1-hiera-* checkpoints (SamGeo2 accepts only SAM 2.0 ids). apply_postprocessing=False as in SamGeo2, so the two differ only in their weights.
  • "sam3": samgeo.SamGeo3(backend="meta", enable_inst_interactivity=True) and predict_inst(box=...) (SAM 3 instance-interactive, i.e. SAM 1/2 style single-object prompts). Needs a CUDA GPU and triton, which the Meta package imports at import time: Linux is the platform Meta supports; Windows works only through the community triton-windows wheel (not verified by agribound; a WARNING is logged); macOS is not supported. Gated weights facebook/sam3 or facebook/sam3.1.
  • "sam3-hf": transformers.Sam3TrackerModel / Sam3TrackerProcessor (promptable visual segmentation, one mask per box). No triton dependency, so it is the SAM 3 option for platforms without triton (Windows without triton-windows, macOS); it imports on macOS, but agribound has not yet run it end to end on any platform because the weights are gated. CUDA recommended. Gated weights facebook/sam3.

Concept-exemplar (PCS) box prompts such as SamGeo3.generate_masks_by_boxes, which segment all objects similar to the box, are never used.

refine_boundaries

refine_boundaries(gdf: GeoDataFrame, raster_path: str, config: AgriboundConfig, **kwargs: Any) -> gpd.GeoDataFrame

Refine field boundaries with box-prompted SAM (see the module docstring).

Parameters:

Name Type Description Default
gdf GeoDataFrame

Field boundaries from a delineation engine.

required
raster_path str

GeoTIFF the polygons were delineated from (not rotated or sheared).

required
config AgriboundConfig

Uses source (RGB band lookup), bands, sam_backend, sam_model (legacy: engine_params["sam_model"]), sam_min_crop_px, sam_crop_padding and device, plus engine_params "sam_rgb_bands" (three 1-based band indices, required for embedding rasters), "sam_window_px" (default 1024), "sam_batch_size" (boxes per decoder call, default 32), "sam_overlaps" ("trim", default, or "keep"; step 6 of the module docstring) and "sam_min_coverage" (number in [0, 1], default :data:DEFAULT_MIN_COVERAGE; 0 disables it; step 7).

required
**kwargs Any

rgb_bands, window_px and batch_size override the corresponding engine_params; predictor supplies an already loaded predictor object with set_image(uint8 HxWx3) and predict_boxes((B, 4) xyxy) -> ((B, H, W) bool, (B,) scores) (plus backend, model_id and device attributes).

{}

Returns:

Type Description
GeoDataFrame

Copy of gdf (same rows, order, index and columns) with geometry replaced where refined, a bool column "agribound:sam_refined" and a float column "agribound:sam_score" (SAM's predicted IoU for refined rows, NaN otherwise). attrs["sam_stats"] holds backend, model, device, n_total, n_refined, n_skipped_small, n_skipped_outside, n_failed, min_crop_px, padding, window_px, max_window_px, batch_size, n_windows, n_windows_decimated, rgb_bands, rgb_source, stretch, errors, overlaps, n_overlap_trimmed, overlap_trimmed_fraction, min_coverage, n_low_coverage (plus multimask_output, mask_selection and, for SAM 2/2.1, apply_postprocessing when SAM ran), with n_total == n_refined + n_skipped_small + n_skipped_outside + n_failed + n_low_coverage. model and device are the configured ones when no polygon is prompted (no model is loaded then). n_skipped_outside counts missing/empty geometries and polygons whose bounding box is not inside the raster (half-pixel tolerance); a polygon that is both outside and small counts as outside. n_failed counts empty masks, masks that lay entirely on other polygons (step 6 of the module docstring) and fields in windows where SAM raised. n_overlap_trimmed counts the masks trimmed so as not to overlap other polygons, and overlap_trimmed_fraction is the share of the refined mask area removed by that trimming (masks counted in n_low_coverage excluded). n_low_coverage counts the polygons that keep their input geometry because their mask covered less than min_coverage of it (step 7).

Raises:

Type Description
ValueError

For a rotated raster, bad band indices, an embedding raster without sam_rgb_bands, or an invalid sam_window_px, sam_batch_size, sam_overlaps or sam_min_coverage.

RuntimeError

If SAM raised for every window that had prompts.

Notes

Masks depend on the compute device: on a Sentinel-2 test crop, SAM 2 (sam2-hiera-tiny) masks computed on Apple MPS overlapped the CPU masks of the same fields with IoU between 0.59 and 0.97. sam_stats records the device.

Source code in agribound/engines/samgeo_engine.py
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
1207
1208
1209
1210
1211
1212
1213
1214
1215
1216
1217
1218
1219
1220
1221
1222
1223
1224
1225
1226
1227
1228
1229
1230
1231
1232
1233
1234
1235
1236
1237
1238
1239
1240
1241
1242
1243
1244
1245
1246
1247
1248
1249
1250
1251
1252
1253
1254
1255
1256
1257
1258
1259
1260
1261
1262
1263
1264
1265
1266
1267
1268
1269
1270
1271
1272
1273
1274
1275
1276
1277
1278
1279
1280
1281
1282
1283
1284
1285
1286
1287
1288
1289
1290
1291
1292
1293
1294
1295
1296
1297
1298
1299
1300
1301
1302
1303
1304
1305
1306
1307
1308
1309
1310
1311
1312
1313
1314
1315
1316
1317
1318
1319
1320
1321
1322
1323
1324
1325
1326
1327
1328
1329
1330
1331
1332
1333
1334
1335
1336
1337
1338
1339
1340
1341
1342
1343
1344
1345
1346
1347
1348
1349
1350
1351
1352
1353
1354
1355
1356
1357
1358
1359
1360
1361
1362
1363
1364
1365
1366
1367
1368
1369
1370
1371
1372
1373
1374
1375
1376
1377
1378
1379
1380
1381
1382
1383
1384
1385
1386
1387
1388
1389
1390
1391
1392
1393
1394
1395
1396
1397
1398
1399
1400
1401
1402
1403
1404
1405
1406
1407
1408
1409
1410
1411
1412
1413
1414
1415
1416
1417
1418
1419
1420
1421
1422
1423
1424
1425
1426
1427
1428
1429
1430
1431
1432
1433
1434
1435
1436
1437
1438
1439
1440
1441
1442
1443
1444
1445
1446
1447
1448
1449
1450
1451
1452
1453
1454
1455
1456
1457
1458
1459
1460
1461
1462
1463
1464
1465
1466
1467
1468
1469
1470
1471
1472
1473
1474
1475
1476
1477
1478
1479
1480
1481
1482
1483
1484
1485
1486
1487
1488
1489
1490
1491
1492
1493
1494
1495
1496
1497
1498
1499
1500
1501
1502
1503
1504
1505
1506
1507
1508
1509
1510
1511
1512
1513
1514
1515
1516
1517
1518
1519
1520
1521
1522
1523
1524
def refine_boundaries(
    gdf: gpd.GeoDataFrame,
    raster_path: str,
    config: AgriboundConfig,
    **kwargs: Any,
) -> gpd.GeoDataFrame:
    """Refine field boundaries with box-prompted SAM (see the module docstring).

    Parameters
    ----------
    gdf : geopandas.GeoDataFrame
        Field boundaries from a delineation engine.
    raster_path : str
        GeoTIFF the polygons were delineated from (not rotated or sheared).
    config : AgriboundConfig
        Uses ``source`` (RGB band lookup), ``bands``, ``sam_backend``,
        ``sam_model`` (legacy: ``engine_params["sam_model"]``),
        ``sam_min_crop_px``, ``sam_crop_padding`` and ``device``, plus
        ``engine_params`` ``"sam_rgb_bands"`` (three 1-based band indices,
        required for embedding rasters), ``"sam_window_px"`` (default 1024),
        ``"sam_batch_size"`` (boxes per decoder call, default 32),
        ``"sam_overlaps"`` (``"trim"``, default, or ``"keep"``; step 6 of the
        module docstring) and ``"sam_min_coverage"`` (number in [0, 1],
        default :data:`DEFAULT_MIN_COVERAGE`; 0 disables it; step 7).
    **kwargs
        ``rgb_bands``, ``window_px`` and ``batch_size`` override the
        corresponding ``engine_params``; ``predictor`` supplies an already
        loaded predictor object with ``set_image(uint8 HxWx3)`` and
        ``predict_boxes((B, 4) xyxy) -> ((B, H, W) bool, (B,) scores)``
        (plus ``backend``, ``model_id`` and ``device`` attributes).

    Returns
    -------
    geopandas.GeoDataFrame
        Copy of *gdf* (same rows, order, index and columns) with geometry
        replaced where refined, a bool column ``"agribound:sam_refined"`` and
        a float column ``"agribound:sam_score"`` (SAM's predicted IoU for
        refined rows, NaN otherwise). ``attrs["sam_stats"]`` holds
        ``backend, model, device, n_total, n_refined, n_skipped_small,
        n_skipped_outside, n_failed, min_crop_px, padding, window_px,
        max_window_px, batch_size, n_windows, n_windows_decimated,
        rgb_bands, rgb_source, stretch, errors, overlaps, n_overlap_trimmed,
        overlap_trimmed_fraction, min_coverage, n_low_coverage`` (plus
        ``multimask_output``, ``mask_selection`` and, for SAM 2/2.1,
        ``apply_postprocessing`` when SAM ran), with ``n_total == n_refined +
        n_skipped_small + n_skipped_outside + n_failed + n_low_coverage``.
        ``model`` and ``device`` are the configured ones when no polygon is
        prompted (no model is loaded then). ``n_skipped_outside`` counts
        missing/empty geometries and polygons whose bounding box is not
        inside the raster (half-pixel tolerance); a polygon that is both
        outside and small counts as outside. ``n_failed`` counts empty masks,
        masks that lay entirely on other polygons (step 6 of the module
        docstring) and fields in windows where SAM raised.
        ``n_overlap_trimmed`` counts the masks trimmed so as not to overlap
        other polygons, and ``overlap_trimmed_fraction`` is the share of the
        refined mask area removed by that trimming (masks counted in
        ``n_low_coverage`` excluded). ``n_low_coverage`` counts the polygons
        that keep their input geometry because their mask covered less than
        ``min_coverage`` of it (step 7).

    Raises
    ------
    ValueError
        For a rotated raster, bad band indices, an embedding raster without
        ``sam_rgb_bands``, or an invalid ``sam_window_px``, ``sam_batch_size``,
        ``sam_overlaps`` or ``sam_min_coverage``.
    RuntimeError
        If SAM raised for every window that had prompts.

    Notes
    -----
    Masks depend on the compute device: on a Sentinel-2 test crop, SAM 2
    (``sam2-hiera-tiny``) masks computed on Apple MPS overlapped the CPU
    masks of the same fields with IoU between 0.59 and 0.97. ``sam_stats``
    records the device.
    """
    import rasterio
    from rasterio.enums import Resampling
    from rasterio.windows import Window

    from agribound.registry import source_value_scale

    backend = config.sam_backend
    min_crop_px = int(config.sam_min_crop_px)
    padding = float(config.sam_crop_padding)
    params = config.engine_params or {}
    window_px = int(kwargs.get("window_px") or params.get("sam_window_px") or DEFAULT_WINDOW_PX)
    batch_size = int(kwargs.get("batch_size") or params.get("sam_batch_size") or DEFAULT_BATCH_SIZE)
    if window_px < 2 * min_crop_px:
        raise ValueError(f"sam_window_px ({window_px}) must be >= 2 * sam_min_crop_px")
    if batch_size < 1:
        raise ValueError(f"sam_batch_size must be >= 1, got {batch_size}")
    overlaps = str(params.get("sam_overlaps", "trim"))
    if overlaps not in SAM_OVERLAP_MODES:
        raise ValueError(f"sam_overlaps must be one of {SAM_OVERLAP_MODES}, got {overlaps!r}")
    min_coverage = _min_coverage(params.get("sam_min_coverage", DEFAULT_MIN_COVERAGE))

    result = gdf.copy()
    result.attrs = dict(gdf.attrs)
    n_total = len(gdf)
    refined_flags = np.zeros(n_total, dtype=bool)
    scores_out = np.full(n_total, np.nan, dtype=np.float64)
    stats: dict[str, Any] = {
        "backend": backend,
        "model": None,
        "device": None,
        "n_total": n_total,
        "n_refined": 0,
        "n_skipped_small": 0,
        "n_skipped_outside": 0,
        "n_failed": 0,
        "min_crop_px": min_crop_px,
        "padding": padding,
        "window_px": window_px,
        "max_window_px": _max_window_px(window_px),
        "batch_size": batch_size,
        "n_windows": 0,
        "n_windows_decimated": 0,
        "rgb_bands": None,
        "stretch": None,
        "errors": [],
        "overlaps": overlaps,
        "n_overlap_trimmed": 0,
        "overlap_trimmed_fraction": 0.0,
        "min_coverage": min_coverage,
        "n_low_coverage": 0,
    }
    predictor = kwargs.get("predictor")
    device = config.resolve_device()
    if predictor is None:
        # The configured model and device, recorded even if nothing is prompted.
        model_id = resolve_sam_model(backend, _configured_model(config))
        stats["model"], stats["device"] = model_id, str(device)
    else:
        stats["backend"] = getattr(predictor, "backend", backend)
        stats["model"] = getattr(predictor, "model_id", None)
        stats["device"] = str(getattr(predictor, "device", device))

    with rasterio.open(raster_path) as src:
        transform = src.transform
        if transform.b != 0 or transform.d != 0:
            raise ValueError(
                f"{raster_path} has a rotated/sheared transform, which SAM refinement does "
                "not support"
            )
        rgb = _rgb_band_indices(config, src.count, kwargs.get("rgb_bands"))
        embedding = source_value_scale(config.source) == "embedding"
        stats["rgb_bands"] = rgb
        stats["rgb_source"] = "embedding dimensions (pseudo-RGB)" if embedding else "imagery"

        raster_crs = src.crs
        if gdf.crs is None:
            logger.warning("refine_boundaries: polygons have no CRS; assuming the raster CRS")
            proj = gdf.set_crs(raster_crs, allow_override=True)
        elif raster_crs is not None and gdf.crs != raster_crs:
            proj = gdf.to_crs(raster_crs)
        else:
            proj = gdf

        pixel_size = (abs(transform.a), abs(transform.e))
        xs = (transform.c, transform.c + transform.a * src.width)
        ys = (transform.f, transform.f + transform.e * src.height)
        raster_bounds = (min(xs), min(ys), max(xs), max(ys))
        geoms = list(proj.geometry)
        present = np.array([g is not None and not g.is_empty for g in geoms], dtype=bool)
        bounds = np.array(
            [g.bounds if ok else (np.nan,) * 4 for g, ok in zip(geoms, present, strict=True)],
            dtype=np.float64,
        ).reshape(n_total, 4)
        inside = np.array(
            [
                ok and _inside_raster(tuple(b), raster_bounds, pixel_size)
                for b, ok in zip(bounds, present, strict=True)
            ],
            dtype=bool,
        )
        # Same decision as is_refinable(..., raster_bounds=raster_bounds).
        refinable = np.array(
            [
                ok and is_refinable(tuple(b), pixel_size, min_crop_px, padding)
                for b, ok in zip(bounds, inside, strict=True)
            ],
            dtype=bool,
        )
        skipped_small = inside & ~refinable

        # Pixel bounds (col0, row0, col1, row1) of each box. The transform has no
        # rotation (checked above), so col = (x - c) / a and row = (y - f) / e;
        # sorting the two corners handles north-up and south-up rasters alike.
        cols_a = (bounds[:, 0] - transform.c) / transform.a
        cols_b = (bounds[:, 2] - transform.c) / transform.a
        rows_a = (bounds[:, 3] - transform.f) / transform.e
        rows_b = (bounds[:, 1] - transform.f) / transform.e
        bounds_px = np.column_stack(
            [
                np.minimum(cols_a, cols_b),
                np.minimum(rows_a, rows_b),
                np.maximum(cols_a, cols_b),
                np.maximum(rows_a, rows_b),
            ]
        )
        windows = _plan_windows(bounds_px, refinable, padding, src.width, src.height, window_px)
        stats["n_skipped_small"] = int(skipped_small.sum())
        stats["n_skipped_outside"] = int((~inside).sum())
        stats["n_windows"] = len(windows)
        stats["n_windows_decimated"] = sum(
            1 for e in windows.values() if e["out_shape"] != (e["window"][3], e["window"][2])
        )

        n_prompts = sum(len(w["items"]) for w in windows.values())
        refined_geoms: dict[int, Any] = {}
        n_failed = 0
        n_window_errors = 0
        if n_prompts:
            lows, highs, passthrough = _stretch_bounds(src, rgb, embedding)
            stats["stretch"] = {
                "method": "none (uint8)" if passthrough else "percentile",
                "percentiles": None if passthrough else list(STRETCH_PERCENTILES),
                "lows": lows,
                "highs": highs,
            }
            if predictor is None:
                logger.info("Loading SAM backend %s (%s) on %s", backend, model_id, device)
                predictor = _load_predictor(backend, model_id, device)
                stats["backend"] = getattr(predictor, "backend", backend)
                stats["model"] = getattr(predictor, "model_id", model_id)
                stats["device"] = str(getattr(predictor, "device", device))
            if backend in ("sam2", "sam2.1"):
                stats["apply_postprocessing"] = False
            stats["multimask_output"] = False
            stats["mask_selection"] = (
                "largest polygon of the mask inside the field's padded box (whole pixels)"
            )

            logger.info(
                "SAM refinement: %d of %d polygons prompted in %d windows "
                "(%d below %d px, %d missing/outside)",
                n_prompts,
                n_total,
                len(windows),
                stats["n_skipped_small"],
                min_crop_px,
                stats["n_skipped_outside"],
            )
            with _quiet_root_info():
                for i_win, entry in enumerate(windows.values()):
                    col, row, w, h = entry["window"]
                    out_h, out_w = entry["out_shape"]
                    items = entry["items"]
                    win = Window(col, row, w, h)
                    done: set[int] = set()
                    try:
                        if (out_h, out_w) == (h, w):
                            data = src.read(rgb, window=win)
                        else:  # oversized field window: decimated read
                            data = src.read(
                                rgb,
                                window=win,
                                out_shape=(len(rgb), out_h, out_w),
                                resampling=Resampling.nearest,
                            )
                        image = _pad_to_square(
                            _apply_stretch(data, lows, highs, src.nodata, passthrough)
                        )
                        predictor.set_image(image)
                        # Raster pixels per output pixel (1 unless decimated).
                        sx, sy = w / out_w, h / out_h
                        for start in range(0, len(items), batch_size):
                            chunk = items[start : start + batch_size]
                            boxes = np.array([b for _, b, _ in chunk], dtype=np.float64)
                            masks, scores = predictor.predict_boxes(boxes)
                            for (pos, _, clip), mask, score in zip(
                                chunk, masks, scores, strict=True
                            ):
                                cx0, cy0, cx1, cy1 = clip
                                poly = _mask_to_polygon(
                                    mask[cy0:cy1, cx0:cx1],
                                    _grid_transform(
                                        transform, col + cx0 * sx, row + cy0 * sy, sx, sy
                                    ),
                                )
                                done.add(pos)
                                if poly is None:
                                    n_failed += 1
                                    continue
                                refined_geoms[pos] = poly
                                scores_out[pos] = float(score)
                    except Exception as exc:  # the rest of this window failed
                        n_window_errors += 1
                        n_failed += sum(1 for p, _, _ in items if p not in done)
                        if len(stats["errors"]) < 5:
                            stats["errors"].append(f"{type(exc).__name__}: {exc}")
                        logger.debug("SAM window %s failed: %s", entry["window"], exc)
                    if (i_win + 1) % 50 == 0:
                        logger.info(
                            "SAM refinement: %d/%d windows done (%d refined)",
                            i_win + 1,
                            len(windows),
                            len(refined_geoms),
                        )

        if n_prompts and n_window_errors == len(windows):
            raise RuntimeError(
                f"SAM refinement failed in all {len(windows)} windows; first error: "
                f"{stats['errors'][0] if stats['errors'] else 'unknown'}"
            )

    low_coverage: list[int] = []
    if refined_geoms and overlaps == "trim":
        # Steps 6 and 7: no refined mask takes area of another polygon, and a trimmed
        # mask covering too little of its input polygon is not used.
        kept, trimmed, removed, low_coverage = _trim_overlaps(
            geoms, refined_geoms, scores_out, min_coverage
        )
        used = set(refined_geoms) - set(low_coverage)
        total_area = float(sum(refined_geoms[p].area for p in used))
        dropped = sorted(used - set(kept))
        for pos in dropped:  # nothing left: keep the input geometry
            scores_out[pos] = np.nan
        n_failed += len(dropped)
        refined_geoms = kept
        stats["n_overlap_trimmed"] = len(trimmed)
        stats["overlap_trimmed_fraction"] = round(removed / total_area, 6) if total_area else 0.0
        if trimmed:
            logger.info(
                "SAM refinement: %d masks trimmed where they overlapped other polygons "
                "(%.1f %% of the refined area; %d dropped)",
                len(trimmed),
                100.0 * stats["overlap_trimmed_fraction"],
                len(dropped),
            )
    elif refined_geoms and min_coverage > 0:  # "keep": step 7 on the masks as SAM drew them
        work = _repaired(geoms)
        for pos in sorted(refined_geoms):
            if _covers_too_little(refined_geoms[pos], work[pos], min_coverage):
                low_coverage.append(pos)
                del refined_geoms[pos]
    scores_out[low_coverage] = np.nan  # these rows keep their input geometry
    stats["n_low_coverage"] = len(low_coverage)
    if low_coverage:
        logger.info(
            "SAM refinement: %d masks covered less than %g %% of their input polygon; those "
            "polygons keep their input geometry (sam_min_coverage=%g)",
            len(low_coverage),
            100.0 * min_coverage,
            min_coverage,
        )

    if refined_geoms:
        positions = sorted(refined_geoms)
        new_geoms = gpd.GeoSeries([refined_geoms[p] for p in positions], crs=raster_crs)
        if gdf.crs is not None and raster_crs is not None and gdf.crs != raster_crs:
            new_geoms = new_geoms.to_crs(gdf.crs)
        geom_col = result.geometry.name
        values = list(result.geometry)
        for p, g in zip(positions, new_geoms, strict=True):
            values[p] = g
        result[geom_col] = gpd.GeoSeries(values, index=result.index, crs=gdf.crs)
        refined_flags[positions] = True

    stats["n_refined"] = int(refined_flags.sum())
    stats["n_failed"] = int(n_failed)
    result[REFINED_COLUMN] = refined_flags
    result[SCORE_COLUMN] = scores_out
    result.attrs["sam_stats"] = stats

    if stats["n_failed"]:
        logger.warning(
            "SAM refinement: %d polygons could not be refined and keep their geometry%s",
            stats["n_failed"],
            f" (first error: {stats['errors'][0]})" if stats["errors"] else " (empty masks)",
        )
    logger.info(
        "SAM refinement (%s, %s): %d refined, %d below %d px, %d missing/outside, %d failed, "
        "%d below sam_min_coverage=%g, of %d polygons",
        stats["backend"],
        stats["model"],
        stats["n_refined"],
        stats["n_skipped_small"],
        min_crop_px,
        stats["n_skipped_outside"],
        stats["n_failed"],
        stats["n_low_coverage"],
        min_coverage,
        n_total,
    )
    return result

crop_window_px

crop_window_px(bounds: tuple[float, float, float, float], pixel_size: tuple[float, float], padding: float) -> tuple[int, int]

Return the padded bounding-box size in whole pixels, as used for gating.

width_px = floor((maxx - minx) / |pixel_size[0]| * (1 + 2 * padding)) and likewise for the height with pixel_size[1] (a tolerance of 1e-6 pixel absorbs floating-point error).

Parameters:

Name Type Description Default
bounds tuple of float

(minx, miny, maxx, maxy) in the raster CRS.

required
pixel_size tuple of float

Pixel width and height in CRS units (sign ignored).

required
padding float

Padding on each side as a fraction of the box size (>= 0).

required

Returns:

Type Description
tuple[int, int]

(width_px, height_px); (0, 0) for non-finite bounds.

Raises:

Type Description
ValueError

If a pixel size is not positive or padding is negative.

Source code in agribound/engines/samgeo_engine.py
def crop_window_px(
    bounds: tuple[float, float, float, float],
    pixel_size: tuple[float, float],
    padding: float,
) -> tuple[int, int]:
    """Return the padded bounding-box size in whole pixels, as used for gating.

    ``width_px = floor((maxx - minx) / |pixel_size[0]| * (1 + 2 * padding))``
    and likewise for the height with ``pixel_size[1]`` (a tolerance of 1e-6
    pixel absorbs floating-point error).

    Parameters
    ----------
    bounds : tuple of float
        ``(minx, miny, maxx, maxy)`` in the raster CRS.
    pixel_size : tuple of float
        Pixel width and height in CRS units (sign ignored).
    padding : float
        Padding on each side as a fraction of the box size (>= 0).

    Returns
    -------
    tuple[int, int]
        ``(width_px, height_px)``; ``(0, 0)`` for non-finite bounds.

    Raises
    ------
    ValueError
        If a pixel size is not positive or *padding* is negative.
    """
    px, py = abs(float(pixel_size[0])), abs(float(pixel_size[1]))
    if not (px > 0 and py > 0):
        raise ValueError(f"pixel_size must be non-zero, got {pixel_size}")
    if padding < 0:
        raise ValueError(f"padding must be >= 0, got {padding}")
    minx, miny, maxx, maxy = (float(v) for v in bounds)
    if not all(math.isfinite(v) for v in (minx, miny, maxx, maxy)):
        return 0, 0
    scale = 1.0 + 2.0 * float(padding)
    width = max(0.0, (maxx - minx) / px * scale)
    height = max(0.0, (maxy - miny) / py * scale)
    return int(math.floor(width + _EPS)), int(math.floor(height + _EPS))

is_refinable

is_refinable(bounds: tuple[float, float, float, float], pixel_size: tuple[float, float], min_crop_px: int = MIN_CROP_SIZE, padding: float = CROP_PADDING, raster_bounds: tuple[float, float, float, float] | None = None) -> bool

Return True if :func:refine_boundaries would prompt SAM with this box.

A polygon is refined only when both sides of its padded bounding box (:func:crop_window_px) are at least min_crop_px pixels. With the defaults (64 px, 15 % padding) the unpadded box must be at least 64 / 1.3 = 49.2 pixels on each side.

:func:refine_boundaries also skips every polygon whose bounding box extends more than :data:EDGE_TOLERANCE_PX (half a pixel) beyond the raster. Without raster_bounds this function assumes the polygon lies inside the raster (normally true for polygons delineated from it); with raster_bounds it applies that test too and then reproduces :func:refine_boundaries exactly.

Parameters:

Name Type Description Default
bounds tuple of float

(minx, miny, maxx, maxy) in the raster CRS.

required
pixel_size tuple of float

Pixel width and height in CRS units.

required
min_crop_px int

Minimum padded side in pixels (config.sam_min_crop_px).

MIN_CROP_SIZE
padding float

Padding fraction (config.sam_crop_padding).

CROP_PADDING
raster_bounds tuple of float or None

Raster extent (minx, miny, maxx, maxy) in its CRS (the order of the two x and the two y values does not matter).

None

Returns:

Type Description
bool
Source code in agribound/engines/samgeo_engine.py
def is_refinable(
    bounds: tuple[float, float, float, float],
    pixel_size: tuple[float, float],
    min_crop_px: int = MIN_CROP_SIZE,
    padding: float = CROP_PADDING,
    raster_bounds: tuple[float, float, float, float] | None = None,
) -> bool:
    """Return *True* if :func:`refine_boundaries` would prompt SAM with this box.

    A polygon is refined only when both sides of its padded bounding box
    (:func:`crop_window_px`) are at least *min_crop_px* pixels. With the
    defaults (64 px, 15 % padding) the unpadded box must be at least
    ``64 / 1.3 = 49.2`` pixels on each side.

    :func:`refine_boundaries` also skips every polygon whose bounding box
    extends more than :data:`EDGE_TOLERANCE_PX` (half a pixel) beyond the
    raster. Without *raster_bounds* this function assumes the polygon lies
    inside the raster (normally true for polygons delineated from it); with
    *raster_bounds* it applies that test too and then reproduces
    :func:`refine_boundaries` exactly.

    Parameters
    ----------
    bounds : tuple of float
        ``(minx, miny, maxx, maxy)`` in the raster CRS.
    pixel_size : tuple of float
        Pixel width and height in CRS units.
    min_crop_px : int
        Minimum padded side in pixels (``config.sam_min_crop_px``).
    padding : float
        Padding fraction (``config.sam_crop_padding``).
    raster_bounds : tuple of float or None
        Raster extent ``(minx, miny, maxx, maxy)`` in its CRS (the order of
        the two x and the two y values does not matter).

    Returns
    -------
    bool
    """
    if raster_bounds is not None and not _inside_raster(bounds, raster_bounds, pixel_size):
        return False
    width, height = crop_window_px(bounds, pixel_size, padding)
    return min(width, height) >= int(min_crop_px)

prefetch

prefetch(config: AgriboundConfig) -> list[str]

Download the weights of the configured SAM backend into the Hugging Face cache.

Afterwards refinement can run with HF_HUB_OFFLINE=1. Files: sam2/sam2.1: the checkpoint named in sam2.build_sam.HF_MODEL_ID_TO_FILENAMES; sam3: config.json, sam3.pt (or sam3.1_multiplex.pt) and the text-encoder vocabulary; sam3-hf: the repository's *.json and *.safetensors files.

Parameters:

Name Type Description Default
config AgriboundConfig

Uses sam_backend and sam_model (or legacy engine_params["sam_model"]).

required

Returns:

Type Description
list[str]

Local paths of the downloaded files (or snapshot directory).

Raises:

Type Description
RuntimeError

If a gated repository refuses access.

Source code in agribound/engines/samgeo_engine.py
def prefetch(config: AgriboundConfig) -> list[str]:
    """Download the weights of the configured SAM backend into the Hugging Face cache.

    Afterwards refinement can run with ``HF_HUB_OFFLINE=1``. Files:
    ``sam2``/``sam2.1``: the checkpoint named in
    ``sam2.build_sam.HF_MODEL_ID_TO_FILENAMES``; ``sam3``: ``config.json``,
    ``sam3.pt`` (or ``sam3.1_multiplex.pt``) and the text-encoder vocabulary;
    ``sam3-hf``: the repository's ``*.json`` and ``*.safetensors`` files.

    Parameters
    ----------
    config : AgriboundConfig
        Uses ``sam_backend`` and ``sam_model`` (or legacy
        ``engine_params["sam_model"]``).

    Returns
    -------
    list[str]
        Local paths of the downloaded files (or snapshot directory).

    Raises
    ------
    RuntimeError
        If a gated repository refuses access.
    """
    from huggingface_hub import hf_hub_download, snapshot_download

    backend = config.sam_backend
    model_id = resolve_sam_model(backend, _configured_model(config))
    try:
        if backend in ("sam2", "sam2.1"):
            try:
                from sam2.build_sam import HF_MODEL_ID_TO_FILENAMES
            except ImportError as exc:
                raise ImportError(
                    "Prefetching SAM 2 weights needs the sam2 package: "
                    "pip install 'agribound[samgeo]'"
                ) from exc
            checkpoint = HF_MODEL_ID_TO_FILENAMES[model_id][1]
            return [hf_hub_download(model_id, checkpoint)]
        if backend == "sam3":
            return [
                hf_hub_download(model_id, "config.json"),
                hf_hub_download(model_id, _SAM3_CHECKPOINTS[model_id]),
                hf_hub_download(_SAM3_BPE[0], _SAM3_BPE[1], repo_type=_SAM3_BPE[2]),
            ]
        return [snapshot_download(model_id, allow_patterns=["*.json", "*.safetensors"])]
    except Exception as exc:
        if _is_gated(exc):
            raise _gated_error(model_id, exc) from exc
        raise

Fine-tuning

finetune

Fine-tuning on reference field boundaries.

:func:fine_tune adapts a fine-tunable engine to user-supplied reference polygons and returns the path of the resulting checkpoint. The pipeline then passes that path to the engine as engine_params["checkpoint_path"].

Engine-specific code lives in private submodules, which the dispatcher imports only when it needs them:

  • _data: training chips, segmentation masks and the train/validation split (_prepare_training_data(raster_path, config, engine) -> Path)
  • _yolo: Delineate-Anything / Ultralytics YOLO (_finetune_yolo(train_dir, config, model_key) -> str)
  • _geoai: GeoAI Mask R-CNN (_finetune_geoai(train_dir, config) -> str)
  • _dinov3: DINOv3 + DPT head via geoai (_finetune_dinov3(train_dir, config) -> str)
  • _prithvi: Prithvi-EO-2.0 via terratorch (_finetune_prithvi(train_dir, config) -> str)
  • _ftw: not dispatched. FTW models are trained with ftw-baselines.

Which engines can be fine-tuned is read from :data:agribound.registry.ENGINE_REGISTRY (fine_tunable). Any other engine raises :class:ValueError with instructions. The engine is never replaced by a different one.

Caching

Each fine-tuning run gets its own directory from :func:agribound._cache.cache_path. The key covers everything :func:agribound._cache.cache_key hashes (study area, source, year/date range, compositing and export settings) plus the engine, the base-model id, a fingerprint of the reference file (resolved path, modification time and size), the number of epochs, the split settings, the seed, config.bands, the engine_params (excluding checkpoint_path and sam_* keys), the engine's default chip-size rule (:func:agribound.engines.finetune._data.chip_size_rule) and, for Delineate-Anything, the version of the training recipe (:data:agribound.engines.finetune._yolo.RECIPE_VERSION), so a checkpoint trained with an earlier recipe is not reused. The trainers receive a copy of the configuration whose :meth:~agribound.config.AgriboundConfig.get_working_dir is that directory. Their chips and checkpoints therefore cannot collide with another run's. finetune_manifest.json in the directory records the checkpoint, and a later call with the same key returns it without retraining.

fine_tune

fine_tune(raster_path: str, config: AgriboundConfig) -> str

Fine-tune the configured engine on reference field boundaries.

Parameters:

Name Type Description Default
raster_path str

Path to the satellite composite GeoTIFF.

required
config AgriboundConfig

Pipeline configuration with reference_boundaries set. The engine is config.engine. config.seed is passed to :func:agribound._repro.seed_everything (Python, NumPy, and torch and Lightning when installed) before the training data is prepared.

required

Returns:

Type Description
str

Absolute path to the fine-tuned model checkpoint (cached or new).

Raises:

Type Description
ValueError

If reference boundaries are not provided, or if config.engine is unknown or not fine-tunable (fine_tunable is False in :data:agribound.registry.ENGINE_REGISTRY, e.g. "ftw", "embedding", "ensemble").

NotImplementedError

If the registry marks the engine as fine-tunable but no trainer is wired up for it here.

RuntimeError

If the trainer does not return an existing checkpoint file.

Source code in agribound/engines/finetune/__init__.py
def fine_tune(
    raster_path: str,
    config: AgriboundConfig,
) -> str:
    """Fine-tune the configured engine on reference field boundaries.

    Parameters
    ----------
    raster_path : str
        Path to the satellite composite GeoTIFF.
    config : AgriboundConfig
        Pipeline configuration with ``reference_boundaries`` set. The engine
        is ``config.engine``. ``config.seed`` is passed to
        :func:`agribound._repro.seed_everything` (Python, NumPy, and torch and
        Lightning when installed) before the training data is prepared.

    Returns
    -------
    str
        Absolute path to the fine-tuned model checkpoint (cached or new).

    Raises
    ------
    ValueError
        If reference boundaries are not provided, or if ``config.engine`` is
        unknown or not fine-tunable (``fine_tunable`` is False in
        :data:`agribound.registry.ENGINE_REGISTRY`, e.g. ``"ftw"``,
        ``"embedding"``, ``"ensemble"``).
    NotImplementedError
        If the registry marks the engine as fine-tunable but no trainer is
        wired up for it here.
    RuntimeError
        If the trainer does not return an existing checkpoint file.
    """
    if config.reference_boundaries is None:
        raise ValueError("reference_boundaries is required for fine-tuning")

    engine = config.engine
    _check_fine_tunable(engine)

    # Derive a model key for per-model checkpoint isolation
    model_key = _get_model_key(engine, config)
    run_dir = _run_dir(config, engine, model_key)

    # Check for cached checkpoint — avoid redundant fine-tuning
    cached = _cached_checkpoint(run_dir)
    if cached is not None:
        logger.info("Using cached fine-tuned checkpoint: %s", cached)
        return cached

    from agribound._repro import seed_everything

    seed_everything(config.seed)

    # Everything the trainers write under get_working_dir() lands in run_dir.
    train_config = config.merged(cache_dir=str(run_dir))

    logger.info("Preparing training data for fine-tuning (%s) in %s", engine, run_dir)
    data_module = importlib.import_module(f"{__name__}._data")
    train_dir = data_module._prepare_training_data(raster_path, train_config, engine)

    module_name, func_name, takes_model_key = _TRAINERS[engine]
    trainer = getattr(importlib.import_module(f"{__name__}.{module_name}"), func_name)
    if takes_model_key:
        checkpoint = trainer(train_dir, train_config, model_key)
    else:
        checkpoint = trainer(train_dir, train_config)

    checkpoint_path = _validate_checkpoint(engine, checkpoint)
    _write_manifest(
        run_dir,
        checkpoint_path,
        {
            "engine": engine,
            "model_key": model_key,
            "reference_fingerprint": _reference_fingerprint(config.reference_boundaries),
            "fine_tune_epochs": config.fine_tune_epochs,
            "fine_tune_split": config.fine_tune_split,
            "fine_tune_val_split": config.fine_tune_val_split,
            "seed": config.seed,
        },
    )
    logger.info("Fine-tuned %s checkpoint: %s", engine, checkpoint_path)
    return str(checkpoint_path)