"""Shared county zonal statistics for raster-based metrics. Area weighting estimates how much of each raster cell lies inside a county by rasterizing the county on a finer grid of sub-cells, then scales each cell by the cosine of its latitude so cells count by their true surface area. Counties that cross the 180th meridian are split so each side is read from its own small raster window. """ from __future__ import annotations import math from typing import Dict, List import numpy as np import shapely from affine import Affine from rasterio.features import rasterize from rasterio.windows import Window, from_bounds from shapely.affinity import translate from shapely.geometry import box, mapping from shapely.geometry.base import BaseGeometry from shapely.ops import unary_union DEFAULT_SUBCELLS = 16 # Largest sub-cell grid rasterized for one county piece; bigger pieces use a coarser grid. SUBCELL_BUDGET = 80_000_000 def split_at_antimeridian(geometry: BaseGeometry) -> List[BaseGeometry]: """Split a lon/lat geometry into pieces that each stay on one side of 180 degrees. A county such as Aleutians West, AK has islands at both +179 and -179 degrees longitude. Its bounding box then spans nearly the whole globe, so each side is returned as its own piece. """ minx, _, maxx, _ = geometry.bounds if maxx - minx <= 180.0: return [geometry] positive: List[BaseGeometry] = [] negative: List[BaseGeometry] = [] for part in getattr(geometry, "geoms", [geometry]): part_minx, _, part_maxx, _ = part.bounds if part_maxx - part_minx > 180.0: # One outline crosses the line: unwrap to 0..360, cut at 180, rewrap. unwrapped = shapely.transform( part, lambda xy: np.column_stack((np.where(xy[:, 0] < 0, xy[:, 0] + 360.0, xy[:, 0]), xy[:, 1])), ) positive.append(unwrapped.intersection(box(0.0, -90.0, 180.0, 90.0))) negative.append(translate(unwrapped.intersection(box(180.0, -90.0, 360.0, 90.0)), xoff=-360.0)) elif part_minx >= 0: positive.append(part) else: negative.append(part) pieces = [unary_union(group) for group in (positive, negative) if group] return [piece for piece in pieces if not piece.is_empty] def geometry_window(source, geometry: BaseGeometry) -> Window: """Return the raster window covering a geometry, padded by one cell on each side.""" window = from_bounds(*geometry.bounds, transform=source.transform) col_start = math.floor(window.col_off) - 1 row_start = math.floor(window.row_off) - 1 col_stop = math.ceil(window.col_off + window.width) + 1 row_stop = math.ceil(window.row_off + window.height) + 1 padded = Window(col_start, row_start, col_stop - col_start, row_stop - row_start) return padded.intersection(Window(0, 0, source.width, source.height)) def subcells_for(shape: tuple[int, int], requested: int) -> int: """Return the finest sub-cell count, up to the request, that fits the budget.""" rows, cols = shape subcells = requested while subcells > 1 and rows * cols * subcells * subcells > SUBCELL_BUDGET: subcells //= 2 return subcells def cell_coverage_fractions( shape: tuple[int, int], transform: Affine, geometry: BaseGeometry, subcells: int ) -> np.ndarray: """Estimate the fraction of each raster cell covered by a geometry.""" rows, cols = shape fine = rasterize( [mapping(geometry)], out_shape=(rows * subcells, cols * subcells), transform=transform * Affine.scale(1.0 / subcells), fill=0, default_value=1, dtype="uint8", ) return fine.reshape(rows, subcells, cols, subcells).mean(axis=(1, 3)) def area_weighted_class_weights( source, geometry: BaseGeometry, *, geographic: bool = True, subcells: int = DEFAULT_SUBCELLS, ) -> Dict[int, float]: """Return the area inside a geometry covered by each value of a categorical raster. Weights are relative surface areas: the fraction of each cell inside the geometry, times cos(latitude) for geographic rasters. Cells equal to 0 or the raster's nodata value are excluded. """ weights: Dict[int, float] = {} pieces = split_at_antimeridian(geometry) if geographic else [geometry] for piece in pieces: window = geometry_window(source, piece) values = source.read(1, window=window, masked=True).filled(0) if source.nodata is not None: values = np.where(values == source.nodata, 0, values) transform = source.window_transform(window) cell_weights = cell_coverage_fractions(values.shape, transform, piece, subcells_for(values.shape, subcells)) if geographic: row_lat = transform.f + (np.arange(values.shape[0]) + 0.5) * transform.e cell_weights = cell_weights * np.cos(np.radians(row_lat))[:, None] counted = (values != 0) & (cell_weights > 0) for value in np.unique(values[counted]): code = int(value) weights[code] = weights.get(code, 0.0) + float(cell_weights[counted & (values == value)].sum()) return weights