"""Operations on the H3 cell graph (CSR neighbours).""" from __future__ import annotations import heapq import numpy as np from scipy import sparse from scipy.sparse import csgraph from scipy.sparse import linalg as splinalg def nbr_mean(g, f): return np.bincount(g.src, weights=f[g.dst], minlength=g.n) / g.counts def nbr_max(g, f): out = np.asarray(f, dtype=np.float64).copy() np.maximum.at(out, g.src, f[g.dst]) return out def nbr_min(g, f): out = np.asarray(f, dtype=np.float64).copy() np.minimum.at(out, g.src, f[g.dst]) return out def diffuse(g, f, iters: int, alpha: float = 0.5, mask=None): f = np.asarray(f, dtype=np.float64) for _ in range(int(iters)): new = (1.0 - alpha) * f + alpha * nbr_mean(g, f) f = new if mask is None else np.where(mask, new, f) return f def laplacian_matrix(g): """Graph Laplacian (per km²): Δf_i ≈ (4/k_i) Σ_j (f_j − f_i)/d_ij² (exact for a regular hex lattice).""" w = 4.0 / (g.counts[g.src] * g.edge_km**2) lap = sparse.csr_matrix((w, (g.src, g.dst)), shape=(g.n, g.n)) return lap - sparse.diags(np.asarray(lap.sum(axis=1)).ravel()) def workers() -> int: """Threads for independent jobs (seasons, components): WORLDGEN_THREADS, default 3. Each job computes exactly what it would alone, so results never depend on it; numpy and scipy's sparse kernels release the GIL.""" import os try: return max(1, int(os.environ.get("WORLDGEN_THREADS", "3"))) except ValueError: return 3 def pmap(fn, items) -> list: """[fn(x) for x in items] on up to workers() threads, in order.""" items = list(items) if workers() <= 1 or len(items) <= 1: return [fn(x) for x in items] from concurrent.futures import ThreadPoolExecutor with ThreadPoolExecutor(min(workers(), len(items))) as ex: return list(ex.map(fn, items)) def smooth_km(g, f, length_km: float, rtol: float = 1e-6, keep: bool = False): """Resolution-independent smoothing: solve (I − L²Δ) s = f (screened Poisson, decay length ≈ L km). keep: keep the matrix on g for the next call (loops over one grid; drop_smooth_cache(g) frees it).""" f = np.asarray(f, dtype=np.float64) scale = float(np.max(np.abs(f))) if f.size else 0.0 if length_km <= 0 or scale == 0.0: return f.copy() f = f / scale # linear system: solve at unit scale (avoids breakdown) cache = g.__dict__.get("_screened", {}) if length_km in cache: A, inv_diag = cache[length_km] else: A = (sparse.identity(g.n, format="csr") - length_km**2 * laplacian_matrix(g)).tocsr() inv_diag = 1.0 / A.diagonal() if keep: g.__dict__.setdefault("_screened", {})[length_km] = (A, inv_diag) s, info = bicgstab_jacobi(A, f, f.copy(), inv_diag, rtol, 5000) if info != 0: raise ValueError(f"smooth_km: solver did not converge (info={info})") return s * scale def _make_bicg_jit(): """bicgstab's vector updates fused into single passes (numba optional; WORLDGEN_NO_JIT=1 turns it off). Each element gets the same operations in the same order as scipy's numpy statements; dot products, norms and sparse products stay the very calls scipy makes, so the iterates — and the answer — are the same floats. (Threads were tried and dropped: on a busy machine they wait more than they work.)""" import os if os.environ.get("WORLDGEN_NO_JIT"): return None try: import numba except ImportError: return None @numba.njit(cache=True, nogil=True) def p_update(p, v, r, omega, beta): # p -= omega*v; p *= beta; p += r for i in range(len(p)): p[i] = (p[i] - omega * v[i]) * beta + r[i] @numba.njit(cache=True, nogil=True) def scale(out, d, x): # out = d * x (the Jacobi preconditioner) for i in range(len(x)): out[i] = d[i] * x[i] @numba.njit(cache=True, nogil=True) def axpy_neg(r, a, v): # r -= a*v for i in range(len(r)): r[i] = r[i] - a * v[i] @numba.njit(cache=True, nogil=True) def x_update(x, alpha, phat, omega, shat): # x += alpha*phat; x += omega*shat for i in range(len(x)): x[i] = (x[i] + alpha * phat[i]) + omega * shat[i] @numba.njit(cache=True, nogil=True) def axpy(x, a, v): # x += a*v for i in range(len(x)): x[i] = x[i] + a * v[i] return p_update, scale, axpy_neg, x_update, axpy _bicg_jit = _make_bicg_jit() def _csr_matvec_into(A): """mv(x, y): y = A @ x, by scipy's own kernel into a reused buffer (A @ x zero-fills a fresh array and calls the same csr_matvec; a fresh 16 MB array costs more in page faults than the product).""" try: from scipy.sparse import _sparsetools fn = _sparsetools.csr_matvec except (ImportError, AttributeError): fn = None M, N = A.shape def mv(x, y): if fn is None: y[:] = A @ x return y.fill(0.0) fn(M, N, A.indptr, A.indices, A.data, x, y) return mv JIT_MIN_N = 50_000 # smaller systems: scipy (thread start-up outweighs the gain; same answer either way) def bicgstab_jacobi(A, b, x0, inv_diag, rtol, maxiter): """scipy.sparse.linalg.bicgstab(A, b, x0=x0, rtol=rtol, maxiter=maxiter, M=diag(inv_diag)) for a CSR matrix and float64 vectors: the same iterates, statement by statement (scipy 1.12+'s pure-Python loop), with the vector updates fused and everything written into preallocated buffers (memory-bound; fresh 16 MB temporaries cost more in page faults than in arithmetic). Falls back to scipy without numba.""" b = np.asarray(b, dtype=np.float64).ravel() if (_bicg_jit is None or A.shape[0] < JIT_MIN_N or not sparse.isspmatrix_csr(A) and not isinstance(A, sparse.csr_array) or A.dtype != np.float64): M = splinalg.LinearOperator(A.shape, matvec=lambda x: inv_diag * x) return splinalg.bicgstab(A, b, x0=x0, rtol=rtol, maxiter=maxiter, M=M) p_update, scale, axpy_neg, x_update, axpy = _bicg_jit mv = _csr_matvec_into(A) inv_diag = np.ascontiguousarray(inv_diag, dtype=np.float64) x = np.array(x0, dtype=np.float64) bnrm2 = np.linalg.norm(b) atol = max(0.0, float(rtol) * float(bnrm2)) if bnrm2 == 0: return b, 0 rhotol = np.finfo(x.dtype.char).eps ** 2 omegatol = rhotol rho_prev, omega, alpha, p, v = None, None, None, None, None r = b - A @ x if x.any() else b.copy() rtilde = r.copy() phat, shat, v, t = np.empty_like(r), np.empty_like(r), np.empty_like(r), np.empty_like(r) for iteration in range(maxiter): if np.linalg.norm(r) < atol: return x, 0 rho = np.dot(rtilde, r) if np.abs(rho) < rhotol: return x, -10 if iteration > 0: if np.abs(omega) < omegatol: return x, -11 beta = (rho / rho_prev) * (alpha / omega) p_update(p, v, r, omega, beta) else: p = r.copy() scale(phat, inv_diag, p) mv(phat, v) rv = np.dot(rtilde, v) if rv == 0: return x, -11 alpha = rho / rv axpy_neg(r, alpha, v) s = r # scipy copies r into s here and reads both unchanged until r -= omega*t if np.linalg.norm(s) < atol: axpy(x, alpha, phat) return x, 0 scale(shat, inv_diag, s) mv(shat, t) omega = np.dot(t, s) / np.dot(t, t) x_update(x, alpha, phat, omega, shat) axpy_neg(r, omega, t) rho_prev = rho return x, maxiter def drop_smooth_cache(g) -> None: """Free the matrices smooth_km keeps on g.""" g.__dict__.pop("_screened", None) def gradient(g, f): """Tangent-plane gradient (f per km): (2/k) Σ_j (f_j − f_i)/d_ij · t_ij.""" w = (f[g.dst] - f[g.src]) / g.edge_km s = np.stack([np.bincount(g.src, weights=w * g.edge_tangents[:, c], minlength=g.n) for c in range(3)], axis=1) return s * (2.0 / g.counts)[:, None] def _csr(g, weights): return sparse.csr_matrix((weights, (g.src, g.dst)), shape=(g.n, g.n)) def nearest_source(g, sources, weights=None): sources = np.asarray(sources, dtype=np.int64) if len(sources) == 0: return np.full(g.n, np.inf), np.full(g.n, -9999, dtype=np.int64) w = g.edge_km if weights is None else weights dist, _, src = csgraph.dijkstra(_csr(g, w), directed=True, indices=sources, min_only=True, return_predecessors=True) return dist, src.astype(np.int64) def distance_to(g, mask): return nearest_source(g, np.flatnonzero(mask))[0] def priority_flood(g, z, sink_mask, eps: float = 0.01): """Barnes (2014) priority-flood + ε: every non-sink cell gets a strictly descending path to a sink.""" sink_mask = np.asarray(sink_mask, dtype=bool) if not sink_mask.any(): raise ValueError("priority_flood: no sink cells") has_open = np.bincount(g.src, weights=(~sink_mask)[g.dst].astype(np.float64), minlength=g.n) > 0 seeds = np.flatnonzero(sink_mask & has_open) z = np.asarray(z, dtype=np.float64) if _flood_jit is not None: return _flood_jit(z.copy(), sink_mask.copy(), np.asarray(g.nbr_ptr, np.int64), np.asarray(g.nbr_idx, np.int64), seeds.astype(np.int64), float(eps)) return _flood_py(z, sink_mask, g.nbr_ptr, g.nbr_idx, seeds, eps) def _flood_py(z, sink_mask, nbr_ptr, nbr_idx, seeds, eps): zf = np.asarray(z, dtype=np.float64).tolist() done = np.asarray(sink_mask).tolist() ptr = np.asarray(nbr_ptr).tolist() idx = np.asarray(nbr_idx).tolist() heap = [(zf[i], i) for i in np.asarray(seeds).tolist()] heapq.heapify(heap) while heap: zc, c = heapq.heappop(heap) for k in range(ptr[c], ptr[c + 1]): n = idx[k] if not done[n]: done[n] = True if zf[n] < zc + eps: zf[n] = zc + eps heapq.heappush(heap, (zf[n], n)) return np.array(zf) def _make_flood_jit(): """_flood_py compiled with numba when it is installed (optional: same heap order, same float steps, same result; WORLDGEN_NO_JIT=1 turns it off).""" import os if os.environ.get("WORLDGEN_NO_JIT"): return None try: import numba except ImportError: return None @numba.njit(cache=True) def flood(zf, done, ptr, idx, seeds, eps): heap = [(zf[i], i) for i in seeds] heapq.heapify(heap) while len(heap): zc, c = heapq.heappop(heap) for k in range(ptr[c], ptr[c + 1]): n = idx[k] if not done[n]: done[n] = True if zf[n] < zc + eps: zf[n] = zc + eps heapq.heappush(heap, (zf[n], n)) return zf return flood _flood_jit = _make_flood_jit() def _make_steep_jit(): """steepest_receivers' per-row maximum, compiled (numba optional, as the flood): the first edge in row order with the largest slope, slopes computed as numpy does. ok=False (empty row, NaN): use the numpy path.""" import os if os.environ.get("WORLDGEN_NO_JIT"): return None try: import numba except ImportError: return None @numba.njit(cache=True) def steep(z, ptr, dst, edge): n = len(ptr) - 1 first = np.empty(n, np.int64) best = np.empty(n) for i in range(n): a, b = ptr[i], ptr[i + 1] if a == b: return first, best, False bk = a bs = (z[i] - z[dst[a]]) / edge[a] if bs != bs: return first, best, False for k in range(a + 1, b): sk = (z[i] - z[dst[k]]) / edge[k] if sk != sk: return first, best, False if sk > bs: bs, bk = sk, k first[i], best[i] = bk, bs return first, best, True return steep _steep_jit = _make_steep_jit() def steepest_receivers(g, z): """Per cell: the neighbour of steepest descent (first in neighbour order on ties), the slope and the edge length; no way down → itself, 0, inf.""" ptr = np.asarray(g.nbr_ptr, np.int64) if _steep_jit is not None and g.n: first, s, ok = _steep_jit(np.asarray(z), ptr, np.asarray(g.nbr_idx, np.int64), g.edge_km) if ok: down = s > 0 recv = np.where(down, g.dst[first], np.arange(g.n)) return recv, np.where(down, s, 0.0), np.where(down, g.edge_km[first], np.inf) slope = (z[g.src] - z[g.dst]) / g.edge_km if g.n == 0 or not (np.diff(ptr) > 0).all() or np.isnan(slope).any(): order = np.lexsort((-slope, g.src)) # general case (empty rows, NaN) first = order[ptr[:-1]] else: # same edge as the stable lexsort, without sorting smax = np.maximum.reduceat(slope, ptr[:-1]) cand = np.flatnonzero(slope == smax[g.src]) rows = g.src[cand] first = cand[np.concatenate([[True], rows[1:] != rows[:-1]])] s = slope[first] down = s > 0 ar = np.arange(g.n) recv = np.where(down, g.dst[first], ar) return recv, np.where(down, s, 0.0), np.where(down, g.edge_km[first], np.inf) def _gather(ptr, arr, sel): counts = ptr[sel + 1] - ptr[sel] tot = int(counts.sum()) if tot == 0: return arr[:0] starts = np.repeat(ptr[sel] - np.concatenate([[0], np.cumsum(counts)[:-1]]), counts) return arr[starts + np.arange(tot)] def receiver_levels(recv): recv = np.asarray(recv, dtype=np.int64) n = len(recv) ar = np.arange(n) root = recv == ar donors = ar[~root] donors = donors[np.argsort(recv[donors], kind="stable")] dptr = np.concatenate([[0], np.cumsum(np.bincount(recv[donors], minlength=n))]) levels, frontier, seen = [], ar[root], 0 while len(frontier): levels.append(frontier) seen += len(frontier) frontier = _gather(dptr, donors, frontier) if seen != n: raise ValueError("receiver_levels: cycle in receivers") return levels def accumulate(recv, levels, w): """Sum w down the receiver tree. Per level only the receivers are touched (a full bincount per level costs levels × cells); the sums are added in donor order from 0, as bincount does: the same floats.""" acc = np.asarray(w, dtype=np.float64).copy() buf = np.zeros(len(acc)) if len(levels) > 1: acc += 0.0 # as the first full-length add did: −0 becomes +0 for lv in reversed(levels[1:]): r = recv[lv] buf[r] = 0.0 np.add.at(buf, r, acc[lv]) acc[r] = acc[r] + buf[r] return acc def _make_components_jit(): """Connected-component labels as scipy's connected_components numbers them — each component (every node not in the mask is one by itself) by the order of its lowest node — by union-find over the edges, with no sparse matrix (numba optional; WORLDGEN_NO_JIT=1 turns it off).""" import os if os.environ.get("WORLDGEN_NO_JIT"): return None try: import numba except ImportError: return None @numba.njit(cache=True, nogil=True) def labels(n, src, dst, mask): parent = np.arange(n) for e in range(src.shape[0]): a, b = src[e], dst[e] if mask[a] and mask[b]: while parent[a] != a: parent[a] = parent[parent[a]] a = parent[a] while parent[b] != b: parent[b] = parent[parent[b]] b = parent[b] if a != b: if a < b: parent[b] = a else: parent[a] = b root_lab = np.full(n, -1, np.int32) out = np.empty(n, np.int32) count = 0 for v in range(n): r = v while parent[r] != r: r = parent[r] if root_lab[r] < 0: root_lab[r] = count count += 1 out[v] = root_lab[r] if mask[v] else -1 return out return labels _components_jit = _make_components_jit() def components(g, mask): mask = np.asarray(mask, dtype=bool) if _components_jit is not None: return _components_jit(g.n, g.src, g.dst, mask) e = mask[g.src] & mask[g.dst] m = sparse.csr_matrix((np.ones(int(e.sum())), (g.src[e], g.dst[e])), shape=(g.n, g.n)) _, lab = csgraph.connected_components(m, directed=False) return np.where(mask, lab, -1) OCEAN_MIN_KM2 = 5.0e6 def ocean_mask(g, z, min_area_km2: float = OCEAN_MIN_KM2): """The connected world ocean plus any separate basin ≥ min_area_km2; smaller interior lows count as land.""" wet = np.asarray(z) <= 0 if not wet.any(): return wet lab = components(g, wet) area = np.bincount(lab[wet], weights=g.area_km2[wet]) keep = area >= min_area_km2 keep[np.argmax(area)] = True return wet & keep[np.maximum(lab, 0)]