betaearth / app.py
github-actions
sync from GitHub 2026-04-21T10:55:13Z
d0f4694
Raw
History Blame Contribute Delete
48 kB
"""BetaEarth Embedding Generator — Streamlit Demo.
Interactive map-based interface for generating dense 10 m geospatial
embeddings. Deployable on HuggingFace Spaces.
Requires betaearth>=0.2.3 (download_s2_cloud_mask, antimeridian check, etc.).
Usage:
cd demo && streamlit run app.py
"""
from __future__ import annotations
from pathlib import Path
import hashlib
import json
import os
import shutil
import tempfile
import threading
import time
import uuid
from datetime import datetime, timezone
from pathlib import Path
import folium
import folium.plugins
import numpy as np
import streamlit as st
from streamlit_folium import st_folium
# ---------------------------------------------------------------------------
# Page config
# ---------------------------------------------------------------------------
st.set_page_config(
page_title="BetaEarth",
page_icon="🥕",
layout="wide",
initial_sidebar_state="expanded",
)
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
RESOLUTION = 10.0
# Max total output size in MB. Overridable via env var for local runs:
# BETAEARTH_MAX_OUTPUT_MB=20000 streamlit run demo/app.py
MAX_OUTPUT_MB = int(os.environ.get("BETAEARTH_MAX_OUTPUT_MB", "3000"))
BYTES_PER_PIXEL = 64 * 4
COMPRESSION_RATIO = 1.0 # embeddings are near-incompressible (L2-normed float32)
HF_DATASET_REPO = "asterisk-labs/betaearth-requests" # Private dataset for request logging
# ---------------------------------------------------------------------------
# Request logging to HuggingFace Dataset
# ---------------------------------------------------------------------------
def _log_request_async(hf_token: str, record: dict) -> None:
"""Fire-and-forget logging of request metadata to HuggingFace Dataset."""
import sys
import tempfile
try:
import pandas as pd
from huggingface_hub import HfApi
df = pd.DataFrame([record])
with tempfile.NamedTemporaryFile(suffix=".parquet", delete=False) as tmp:
df.to_parquet(tmp.name, index=False)
tmp_path = tmp.name
timestamp_clean = record["timestamp"].replace(":", "-").replace(".", "-")
file_name = f"requests/{timestamp_clean}.parquet"
api = HfApi(token=hf_token)
api.upload_file(
path_or_fileobj=tmp_path,
path_in_repo=file_name,
repo_id=HF_DATASET_REPO,
repo_type="dataset",
commit_message=f"Log request {record['timestamp']}",
)
Path(tmp_path).unlink()
print(f"[log_request] Uploaded {file_name}", file=sys.stderr, flush=True)
except Exception as e:
# Print to stderr so failures are visible in Space logs but don't interrupt UX
print(f"[log_request] FAILED: {type(e).__name__}: {e}", file=sys.stderr, flush=True)
def log_request(
bbox: tuple[float, float, float, float],
area_km2: float,
years: list[int],
time_mode: str,
custom_dates: tuple[str, str] | None,
save_per_timestamp: bool,
save_per_timestamp_input: bool,
) -> None:
"""Log request metadata asynchronously (fire-and-forget)."""
import sys
# Read token from env (HF Spaces exposes secrets as env vars).
# st.secrets is read in main thread to avoid ScriptRunContext issues in threads.
hf_token = os.environ.get("HF_TOKEN")
if not hf_token:
try:
hf_token = st.secrets.get("HF_TOKEN")
except Exception:
hf_token = None
if not hf_token:
print("[log_request] No HF_TOKEN available, skipping log", file=sys.stderr, flush=True)
return
timestamp = datetime.now(timezone.utc).isoformat().replace("+00:00", "Z")
record = {
"timestamp": timestamp,
"bbox_w": bbox[0],
"bbox_s": bbox[1],
"bbox_e": bbox[2],
"bbox_n": bbox[3],
"area_km2": float(area_km2),
"years": years,
"time_mode": time_mode,
"custom_start": custom_dates[0] if custom_dates else None,
"custom_end": custom_dates[1] if custom_dates else None,
"save_per_timestamp": bool(save_per_timestamp),
"save_per_timestamp_input": bool(save_per_timestamp_input),
}
thread = threading.Thread(
target=_log_request_async,
args=(hf_token, record),
daemon=True,
)
thread.start()
# ---------------------------------------------------------------------------
# Styling
# ---------------------------------------------------------------------------
st.markdown("""
<style>
@font-face {
font-family: 'League Spartan';
src: url('https://raw.githubusercontent.com/asterisk-labs/asterisk-labs.github.io/main/assets/fonts/league-spartan-v11-latin_latin-ext-regular.woff2') format('woff2');
font-weight: 400;
font-style: normal;
}
@font-face {
font-family: 'League Spartan';
src: url('https://raw.githubusercontent.com/asterisk-labs/asterisk-labs.github.io/main/assets/fonts/league-spartan-v11-latin_latin-ext-500.woff2') format('woff2');
font-weight: 500;
font-style: normal;
}
@font-face {
font-family: 'League Spartan';
src: url('https://raw.githubusercontent.com/asterisk-labs/asterisk-labs.github.io/main/assets/fonts/league-spartan-v11-latin_latin-ext-600.woff2') format('woff2');
font-weight: 600;
font-style: normal;
}
@font-face {
font-family: 'League Spartan';
src: url('https://raw.githubusercontent.com/asterisk-labs/asterisk-labs.github.io/main/assets/fonts/league-spartan-v11-latin_latin-ext-700.woff2') format('woff2');
font-weight: 700;
font-style: normal;
}
/* Base — force League Spartan everywhere */
html, body, .stApp, [data-testid="stAppViewContainer"],
.stMarkdown, .stMarkdown p, .stMarkdown span, .stMarkdown li,
.stRadio label, .stCheckbox label, .stSlider label,
.stSelectSlider label, .stDateInput label, .stTextInput label,
[data-testid="stMetricLabel"], [data-testid="stMetricValue"], [data-testid="stMetricDelta"],
[data-baseweb="select"], [data-baseweb="input"], [data-baseweb="radio"] label,
[data-baseweb="toggle"] + div,
button, input, select, textarea,
.stCaption, .stAlert, .stToast,
h1, h2, h3, h4, h5, h6, p, span, label, div {
font-family: 'League Spartan', sans-serif !important;
}
html, body, .stApp, [data-testid="stAppViewContainer"] {
background: #ffffff !important;
}
.block-container { padding: 0 2rem 2rem 2rem !important; max-width: 100% !important; }
header[data-testid="stHeader"] { display: none !important; }
/* Links — use the brand purple everywhere */
a, a:link, a:visited,
.stMarkdown a, .stCaption a {
color: #492ae8 !important;
text-decoration: none !important;
}
a:hover, .stMarkdown a:hover, .stCaption a:hover {
color: #3a1fd0 !important;
text-decoration: underline !important;
}
/* Sidebar */
[data-testid="stSidebar"] {
background: #fafafa !important;
border-right: 1px solid #e8e8e8 !important;
}
[data-testid="collapsedControl"],
button[kind="headerNoPadding"] {
display: none !important;
}
[data-testid="stSidebar"] h1 {
font-family: 'League Spartan', sans-serif !important;
font-size: 1.6rem !important;
font-weight: 600 !important;
color: #1a1a1a !important;
margin-bottom: 0 !important;
}
[data-testid="stSidebar"] .stCaption {
color: #888 !important;
font-size: 0.8rem !important;
}
/* Mobile: inline the sidebar above the map. Desktop layout is unchanged.
Streamlit wraps content in position:absolute + fixed-height containers
which collapse the sidebar under a grid/flex child-sized layout, so we
first neutralise the fixed positioning, then stack with flex. */
@media (max-width: 768px) {
/* Break out of fixed absolute containers so the page flows vertically */
[data-testid="stApp"],
[data-testid="stAppViewContainer"] {
position: static !important;
height: auto !important;
min-height: 100vh !important;
max-height: none !important;
overflow: visible !important;
display: flex !important;
flex-direction: column !important;
}
/* Sidebar: sized to its content, at the top */
[data-testid="stSidebar"] {
order: 0 !important;
position: relative !important;
width: 100% !important;
min-width: 0 !important;
max-width: 100% !important;
height: auto !important;
min-height: fit-content !important;
max-height: none !important;
transform: none !important;
top: auto !important; left: auto !important;
right: auto !important; bottom: auto !important;
margin: 0 !important;
box-shadow: none !important;
border-right: none !important;
border-bottom: 1px solid #e8e8e8 !important;
z-index: auto !important;
flex: 0 0 auto !important;
}
/* Sidebar descendants: allow natural height */
[data-testid="stSidebar"] > div,
[data-testid="stSidebar"] > section,
[data-testid="stSidebar"] [data-testid="stSidebarContent"],
[data-testid="stSidebar"] [data-testid="stSidebarUserContent"] {
position: relative !important;
height: auto !important;
min-height: 0 !important;
max-height: none !important;
transform: none !important;
overflow: visible !important;
}
/* Main content wrapper (unnamed div sibling of stSidebar) + stMain */
[data-testid="stAppViewContainer"] > div,
[data-testid="stMain"],
section.main,
.main {
order: 1 !important;
position: relative !important;
width: 100% !important;
min-width: 0 !important;
max-width: 100% !important;
margin-left: 0 !important;
margin-right: 0 !important;
left: auto !important;
top: auto !important;
z-index: auto !important;
flex: 1 1 auto !important;
}
[data-testid="stSidebar"] .block-container {
padding: 1rem !important;
}
.main .block-container,
[data-testid="stMain"] .block-container {
padding: 0 1rem 1rem 1rem !important;
}
}
/* Metrics — yellow accent border */
[data-testid="stMetric"] {
background: #fffef5;
border-radius: 10px;
padding: 12px 16px;
border: 1px solid #f7cc09;
border-left: 4px solid #f7cc09;
}
[data-testid="stMetricValue"] {
font-size: 1.3rem !important;
font-weight: 700 !important;
color: #492ae8 !important;
}
[data-testid="stMetricDelta"] svg {
fill: #c4ffc2 !important;
}
/* Buttons — primary: blue, download: green */
.stButton > button[kind="primary"] {
background: #492ae8 !important;
color: white !important;
border: none !important;
border-radius: 10px !important;
padding: 12px 24px !important;
font-weight: 700 !important;
font-size: 1rem !important;
transition: all 0.2s !important;
}
.stButton > button[kind="primary"]:hover {
background: #3a1fd0 !important;
transform: translateY(-1px) !important;
box-shadow: 0 4px 12px rgba(73,42,232,0.3) !important;
}
.stDownloadButton > button {
background: #c4ffc2 !important;
color: #1a1a1a !important;
border: 2px solid #492ae8 !important;
border-radius: 10px !important;
font-weight: 700 !important;
transition: all 0.2s !important;
}
.stDownloadButton > button:hover {
background: #a8f0a6 !important;
transform: translateY(-1px) !important;
}
/* Progress bar — only target the fill element, not the text */
.stProgress [role="progressbar"] > div:first-child {
background: linear-gradient(90deg, #492ae8, #c4ffc2) !important;
border-radius: 4px !important;
}
.stProgress {
margin: 8px 0 !important;
}
.stProgress p {
font-size: 0.85rem !important;
color: #1a1a1a !important;
margin-bottom: 4px !important;
}
/* Slider track fill — blue */
.stSlider div[data-baseweb="slider"] div[role="progressbar"] {
background: #492ae8 !important;
}
/* Toggle track — blue when on */
[data-baseweb="toggle"] > div:first-child {
background-color: #492ae8 !important;
}
/* Info boxes — soft purple */
.stAlert [data-baseweb="notification"] {
background: #f0eeff !important;
border-left-color: #492ae8 !important;
}
/* Leaflet draw toolbar — larger icons */
.leaflet-draw-toolbar a {
width: 36px !important;
height: 36px !important;
line-height: 36px !important;
background-size: 360px 30px !important;
}
.leaflet-draw-toolbar {
margin-top: 6px !important;
}
/* Gallery images */
[data-testid="stImage"] {
border-radius: 8px;
overflow: hidden;
box-shadow: 0 2px 8px rgba(0,0,0,0.08);
}
/* Dividers */
hr { border-color: #eee !important; }
/* Badge / tag style for bbox coords */
code {
background: #f0eeff !important;
color: #492ae8 !important;
border-radius: 4px !important;
font-size: 0.75rem !important;
}
/* Section headers below map */
h3 { color: #492ae8 !important; font-weight: 700 !important; }
/* Bottom panel — dark console look */
[data-testid="stBottomBlockContainer"],
.main > div:last-child {
background: #1a1a2e !important;
border-top: 3px solid #492ae8 !important;
padding: 16px 24px !important;
}
/* Console-style text in bottom area */
.console-panel {
background: #1a1a2e;
border-top: 3px solid #492ae8;
border-radius: 0;
padding: 20px 28px;
margin: 0 -1rem;
color: #e0e0e0;
font-family: 'JetBrains Mono', 'Fira Code', monospace;
font-size: 0.85rem;
}
.console-panel h3 {
color: #c4ffc2 !important;
font-family: 'League Spartan', sans-serif !important;
font-size: 1.1rem !important;
margin-bottom: 12px !important;
}
.console-panel .stProgress [role="progressbar"] > div:first-child {
background: linear-gradient(90deg, #492ae8, #c4ffc2) !important;
}
.console-panel p, .console-panel span {
color: #e0e0e0 !important;
}
.console-panel code {
background: #2a2a4a !important;
color: #f7cc09 !important;
}
</style>
""", unsafe_allow_html=True)
# ---------------------------------------------------------------------------
# Size estimation
# ---------------------------------------------------------------------------
MIN_SIDE_KM = 3.2 # ~320 px at 10 m — safe multiple of 32 and > tile_size 224
def estimate_size(
bbox,
n_scenes=28,
n_years=1,
save_per_timestamp=True,
save_per_timestamp_input=False,
):
"""Return (width_km, height_km, total_mb). Factors in per-timestamp toggles."""
w, s, e, n = bbox
width_km = (e - w) * 111 * abs(np.cos(np.radians((s + n) / 2)))
height_km = (n - s) * 111
n_pixels = int(width_km * 1000 / RESOLUTION) * int(height_km * 1000 / RESOLUTION)
emb_bytes = n_pixels * BYTES_PER_PIXEL # 64 bands × f32
s2_bytes = n_pixels * 9 * 4 # 9 bands × f32
s1_bytes = n_pixels * 2 * 4 # 2 bands × f32
dem_bytes = n_pixels * 4 # 1 band × f32
# Rough split: ~24 S2 + ~4 S1 per year at max quota
s2_per_year = 24 * (n_scenes / 28)
s1_per_year = 4 * (n_scenes / 28)
total = emb_bytes * n_years # 1 annual .tif per year
if save_per_timestamp:
total += emb_bytes * (s2_per_year + s1_per_year) * n_years
if save_per_timestamp_input:
total += s2_bytes * s2_per_year * n_years
total += s1_bytes * s1_per_year * n_years
total += dem_bytes # one-time
return round(width_km, 1), round(height_km, 1), round(total / 1e6, 1)
def expand_to_min(bbox, min_km=MIN_SIDE_KM):
"""Return (was_padded, required_pad_px_x, required_pad_px_y). Pixels of padding
needed per side (east/west / north/south) so each side >= min_km. Padding is
computed in pixel space so it can be applied directly to the user grid's UTM
bounds — this avoids the UTM-conformal-distortion bug of re-projecting a
lon/lat-expanded bbox."""
w, s, e, n = bbox
lat_mid = (s + n) / 2
width_km = (e - w) * 111 * abs(np.cos(np.radians(lat_mid)))
height_km = (n - s) * 111
min_px = int(np.ceil(min_km * 1000 / RESOLUTION))
cur_px_x = int(np.ceil(width_km * 1000 / RESOLUTION))
cur_px_y = int(np.ceil(height_km * 1000 / RESOLUTION))
pad_px_x = max(0, (min_px - cur_px_x + 1) // 2)
pad_px_y = max(0, (min_px - cur_px_y + 1) // 2)
was_padded = pad_px_x > 0 or pad_px_y > 0
return was_padded, pad_px_x, pad_px_y
def pad_grid_from_user(user_grid, pad_px_x, pad_px_y):
"""Build a padded grid whose UTM bounds strictly contain user_grid, avoiding
UTM conformal distortion issues. The padded grid shares user_grid's CRS and
is aligned to its pixel lattice."""
import rasterio.transform
if pad_px_x == 0 and pad_px_y == 0:
return user_grid
h, w = user_grid["shape"]
xoff = user_grid["transform"].xoff - pad_px_x * RESOLUTION
yoff = user_grid["transform"].yoff + pad_px_y * RESOLUTION # yoff is top, +north
new_w = w + 2 * pad_px_x
new_h = h + 2 * pad_px_y
new_transform = rasterio.transform.from_origin(xoff, yoff, RESOLUTION, RESOLUTION)
x1 = xoff + new_w * RESOLUTION
y0 = yoff - new_h * RESOLUTION
return {
"bbox_4326": user_grid["bbox_4326"], # informational only
"epsg": user_grid["epsg"],
"crs": user_grid["crs"],
"transform": new_transform,
"bounds": (xoff, y0, x1, yoff),
"shape": (new_h, new_w),
}
def crop_to_user(arr, pad_grid, user_grid, channel_axis=-1):
"""Crop an array from the padded grid down to the user's grid."""
if pad_grid is user_grid:
return arr
col_off = round((user_grid["transform"].xoff - pad_grid["transform"].xoff) / RESOLUTION)
row_off = round((pad_grid["transform"].yoff - user_grid["transform"].yoff) / RESOLUTION)
out_h, out_w = user_grid["shape"]
# Defensive clamp — with pad_grid_from_user these should never go negative,
# but old call sites could still trigger the conformal-distortion bug.
if col_off < 0 or row_off < 0:
raise ValueError(
f"pad_grid does not contain user_grid: col_off={col_off}, row_off={row_off}"
)
if channel_axis == 0:
return arr[:, row_off:row_off + out_h, col_off:col_off + out_w]
return arr[row_off:row_off + out_h, col_off:col_off + out_w, ...]
# ---------------------------------------------------------------------------
# Sidebar
# ---------------------------------------------------------------------------
with st.sidebar:
st.image(
"https://github.com/user-attachments/assets/91b2a46c-a142-4eed-99ce-a1b692178146",
use_container_width=True,
)
st.markdown(
"Open-source emulator of the AlphaEarth Foundations (AEF) embedding field. "
"Generate 10 m dense 64-band embeddings from Sentinel-2 + Sentinel-1."
)
st.markdown(
"[📦 GitHub](https://github.com/asterisk-labs/beta-earth) · "
"[🤗 Model Weights](https://huggingface.co/collections/asterisk-labs/beta-earth)"
)
st.divider()
# Default to a narrow custom window so the first-run experience on the
# free CPU Space stays fast (~6 scenes max vs ~28 for a full annual mean).
time_mode = st.radio(
"Time range", ["Custom dates", "Annual"], horizontal=True,
help="Custom = quick ~1-month test. Annual = full-year mosaic (slower on CPU).",
)
if time_mode == "Annual":
year_range = st.slider("Years", 2017, 2025, (2023, 2023))
years = list(range(year_range[0], year_range[1] + 1))
custom_dates = None
else:
from datetime import date
col_d1, col_d2 = st.columns(2)
with col_d1:
start_date = st.date_input("Start", date(2023, 7, 1), min_value=date(2017, 1, 1))
with col_d2:
end_date = st.date_input("End", date(2023, 7, 31), max_value=date(2025, 12, 31))
years = sorted(set(range(start_date.year, end_date.year + 1)))
custom_dates = (str(start_date), str(end_date))
min_coverage = st.slider("Min scene coverage (%)", 50, 100, 100, step=5)
max_cloud = st.slider("Max cloud cover (%)", 5, 50, 20, step=5)
save_per_timestamp = st.toggle("Save per-timestamp embeddings", value=True,
help="Disable to only output the annual average (much smaller download)")
save_per_timestamp_input = st.toggle("Save per-timestamp input", value=False,
help="Also save the raw S2/S1 scene data used for each timestamp (large: adds raw bands per scene)")
st.divider()
# Bbox display
if "bbox" in st.session_state and st.session_state.bbox:
bbox = st.session_state.bbox
n_years = len(years)
w_km, h_km, total_mb = estimate_size(
bbox, n_years=n_years,
save_per_timestamp=save_per_timestamp,
save_per_timestamp_input=save_per_timestamp_input,
)
st.metric("Area", f"{w_km} × {h_km} km")
label = f"{total_mb} MB" + (f" ({n_years} yr)" if n_years > 1 else "")
if total_mb <= MAX_OUTPUT_MB:
st.metric("Est. output", label, delta="OK", delta_color="normal")
else:
st.metric("Est. output", label, delta=f">{MAX_OUTPUT_MB} MB!", delta_color="inverse")
was_padded, _, _ = expand_to_min(bbox)
if was_padded:
st.info(f"Small area — will be padded internally to {MIN_SIDE_KM} km per side, output cropped back.")
st.code(f"W={bbox[0]:.4f}\nS={bbox[1]:.4f}\nE={bbox[2]:.4f}\nN={bbox[3]:.4f}", language=None)
if st.button("🗑 Clear / redraw", use_container_width=True):
st.session_state.pop("bbox", None)
st.session_state.pop("results", None)
st.rerun()
else:
st.info("Draw a rectangle on the map")
st.divider()
generate_btn = st.button(
"🚀 Generate Embeddings",
type="primary",
use_container_width=True,
disabled="bbox" not in st.session_state or not st.session_state.bbox,
)
if "results" in st.session_state and st.session_state.results:
st.divider()
st.caption(st.session_state.results["summary"])
st.divider()
st.caption(
"[GitHub](https://github.com/asterisk-labs/beta-earth) · "
"[Google Satellite Embedding (AEF)](https://developers.google.com/earth-engine/datasets/catalog/GOOGLE_SATELLITE_EMBEDDING_V1_ANNUAL)"
)
# ---------------------------------------------------------------------------
# Map
# ---------------------------------------------------------------------------
m = folium.Map(
location=[20.0, 0.0],
zoom_start=2,
tiles=None,
control_scale=True,
)
# Satellite (togglable)
folium.TileLayer(
tiles="https://server.arcgisonline.com/ArcGIS/rest/services/World_Imagery/MapServer/tile/{z}/{y}/{x}",
attr="Esri",
name="Satellite",
overlay=False,
).add_to(m)
# CartoDB Positron — clean minimal basemap (no API key needed)
folium.TileLayer(
tiles="https://{s}.basemaps.cartocdn.com/light_all/{z}/{x}/{y}{r}.png",
attr="CartoDB",
name="Light",
overlay=False,
show=True,
).add_to(m)
# Draw control for bbox
folium.plugins.Draw(
draw_options={
"polyline": False, "polygon": False, "circle": False,
"circlemarker": False, "marker": False,
"rectangle": {
"showArea": True, # live area tooltip while dragging
"metric": True, # m² / km²
"shapeOptions": {"color": "#492ae8", "weight": 3, "fillOpacity": 0.1},
},
},
edit_options={"edit": False},
).add_to(m)
# Compute max area (km²) for current settings so we can turn the rectangle
# red *during* drawing when the user crosses the output cap. Matches estimate_size().
_N_YEARS = len(years)
_N_PIXELS_PER_KM2 = (1000 / RESOLUTION) ** 2 # = 10_000
_MB_PER_KM2 = 2.56 * _N_YEARS # base: 1 annual 64-band tif per year
if save_per_timestamp:
_MB_PER_KM2 += 2.56 * 28 * _N_YEARS # ~28 per-timestamp embedding tifs
if save_per_timestamp_input:
_MB_PER_KM2 += (0.36 * 24 + 0.08 * 4) * _N_YEARS # raw S2 + S1 per timestamp
_MAX_AREA_KM2 = MAX_OUTPUT_MB / max(_MB_PER_KM2, 1e-3)
# Inject JS: (1) format live draw-tooltip as km² (Leaflet defaults to ha < 1 km²),
# (2) monkey-patch L.Draw.Rectangle._drawShape so the rectangle outline
# turns red live when area > _MAX_AREA_KM2.
from branca.element import MacroElement as _MacroElement
from jinja2 import Template as _Template
class _LiveDrawFeedback(_MacroElement):
_template = _Template("""
{% macro script(this, kwargs) %}
(function () {
// Force km² display in the Leaflet.Draw area tooltip
if (L.GeometryUtil && L.GeometryUtil.readableArea) {
L.GeometryUtil.readableArea = function (area, isMetric, precision) {
var km2 = area / 1e6;
return km2.toFixed(km2 >= 10 ? 1 : 2) + ' km²';
};
}
// Live colour feedback on the rectangle being drawn
var MAX_KM2 = {{ this.max_km2 }};
var OK_COLOR = '#492ae8';
var OVER_COLOR = '#dc2626';
if (L.Draw && L.Draw.Rectangle) {
var orig = L.Draw.Rectangle.prototype._drawShape;
L.Draw.Rectangle.prototype._drawShape = function (latlng) {
orig.call(this, latlng);
if (this._shape) {
var b = this._shape.getBounds();
var w_m = b.getSouthWest().distanceTo(b.getSouthEast());
var h_m = b.getSouthWest().distanceTo(b.getNorthWest());
var km2 = (w_m * h_m) / 1e6;
this._shape.setStyle({color: km2 > MAX_KM2 ? OVER_COLOR : OK_COLOR});
}
};
}
})();
{% endmacro %}
""")
def __init__(self, max_km2):
super().__init__()
self.max_km2 = max_km2
m.add_child(_LiveDrawFeedback(max_km2=_MAX_AREA_KM2))
# Persist the drawn bbox as a visible rectangle + fit view, OR auto-activate
# the draw tool if no bbox has been set yet. Rectangle turns red if the
# current estimate would exceed the output cap.
if "bbox" in st.session_state and st.session_state.bbox:
sbbox = st.session_state.bbox
_, _, _est_mb = estimate_size(
sbbox, n_years=len(years),
save_per_timestamp=save_per_timestamp,
save_per_timestamp_input=save_per_timestamp_input,
)
_rect_color = "#dc2626" if _est_mb > MAX_OUTPUT_MB else "#492ae8"
folium.Rectangle(
bounds=[[sbbox[1], sbbox[0]], [sbbox[3], sbbox[2]]],
color=_rect_color, weight=3, fill=True, fill_opacity=0.1,
).add_to(m)
m.fit_bounds([[sbbox[1], sbbox[0]], [sbbox[3], sbbox[2]]])
else:
from branca.element import MacroElement
from jinja2 import Template
class AutoDrawRectangle(MacroElement):
_template = Template("""
{% macro script(this, kwargs) %}
setTimeout(function() {
var btn = document.querySelector('.leaflet-draw-draw-rectangle');
if (btn) btn.click();
}, 300);
{% endmacro %}
""")
m.add_child(AutoDrawRectangle())
# Add PCA overlay if results exist
if "results" in st.session_state and st.session_state.results:
res = st.session_state.results
opacity = st.session_state.get("opacity", 0.7)
bbox = res["bbox"]
previews = res.get("previews", [])
# Find the selected preview frame. previews is a list of (bytes, label).
selected_label = st.session_state.get("preview_frame")
preview_bytes = None
if previews:
preview_map = {label: data for data, label in previews}
preview_bytes = preview_map.get(selected_label, previews[0][0])
if preview_bytes:
import base64
img_data = base64.b64encode(preview_bytes).decode()
img_url = f"data:image/png;base64,{img_data}"
folium.raster_layers.ImageOverlay(
image=img_url,
bounds=[[bbox[1], bbox[0]], [bbox[3], bbox[2]]],
opacity=opacity,
name="BetaEarth PCA",
).add_to(m)
# Zoom to bbox
m.fit_bounds([[bbox[1], bbox[0]], [bbox[3], bbox[2]]])
folium.LayerControl().add_to(m)
# Status slot ABOVE the map so progress + completion banner are always visible
# without scrolling, even if the map occupies the full viewport height.
status_placeholder = st.empty()
if "results" in st.session_state and st.session_state.results:
with status_placeholder.container():
st.success(
"✓ Embeddings generated — scroll below the map for PCA previews and the ZIP download."
)
# Render map (full width). Height kept modest so progress + results don't
# get hidden below the fold on typical laptop screens.
map_data = st_folium(m, height=650, use_container_width=True, returned_objects=["all_drawings"])
# Extract bbox from drawn rectangle
if map_data and map_data.get("all_drawings"):
drawings = map_data["all_drawings"]
if drawings:
last = drawings[-1]
if last["geometry"]["type"] == "Polygon":
coords = last["geometry"]["coordinates"][0]
# Leaflet lets users pan past ±180° and draw on "repeated" world
# copies, producing coordinates outside the valid range. Normalise
# each longitude into [-180, 180] before assembling a bbox.
def _wrap_lon(x: float) -> float:
return ((x + 180.0) % 360.0) - 180.0
lons_raw = [c[0] for c in coords]
lons = [_wrap_lon(x) for x in lons_raw]
lats = [max(-90.0, min(90.0, c[1])) for c in coords]
w, e = min(lons), max(lons)
# Detect an antimeridian-crossing draw: the raw span is small but
# after wrapping the bbox looks huge (covers most of the globe).
raw_span = max(lons_raw) - min(lons_raw)
wrapped_span = e - w
if raw_span < wrapped_span - 1e-6:
st.warning(
"Your bbox crosses the antimeridian (±180°). BetaEarth doesn't "
"support antimeridian-crossing AOIs yet — please redraw on a "
"single side of the line."
)
else:
new_bbox = (w, min(lats), e, max(lats))
if st.session_state.get("bbox") != new_bbox:
st.session_state.bbox = new_bbox
st.rerun() # rerun so sidebar picks up the new bbox immediately
# ---------------------------------------------------------------------------
# Generation
# ---------------------------------------------------------------------------
if generate_btn and "bbox" in st.session_state and st.session_state.bbox:
bbox = st.session_state.bbox
n_years = len(years)
w_km, h_km, total_mb = estimate_size(
bbox, n_years=n_years,
save_per_timestamp=save_per_timestamp,
save_per_timestamp_input=save_per_timestamp_input,
)
if total_mb > MAX_OUTPUT_MB:
@st.dialog("Area too large for the public demo")
def _too_large_dialog():
st.markdown(
f"The estimated output is **{total_mb:.0f} MB**, which exceeds "
f"the **{MAX_OUTPUT_MB} MB** cap set for this free public demo."
)
st.markdown(
"This demo runs on a free HuggingFace Space with limited CPU and "
"memory, so we keep the ceiling low to keep it usable for everyone."
)
st.markdown("**Options to keep going:**")
st.markdown(
"- Pick a **smaller region** on the map, or reduce the number of years / toggle off per-timestamp saves.\n"
"- Generate **independently on your own machine** — the full pipeline "
"ships in our GitHub repo ([asterisk-labs/beta-earth](https://github.com/asterisk-labs/beta-earth)), "
"including a `betaearth-generate` CLI and a local version of this app "
"with a configurable output cap (`BETAEARTH_MAX_OUTPUT_MB`)."
)
if st.button("OK", use_container_width=True):
st.rerun()
_too_large_dialog()
else:
# Log request metadata up-front so failed generations are also captured
log_request(
bbox=bbox,
area_km2=w_km * h_km,
years=years,
time_mode=time_mode,
custom_dates=custom_dates,
save_per_timestamp=save_per_timestamp,
save_per_timestamp_input=save_per_timestamp_input,
)
# Progress bar rendered into the slot ABOVE the map so it's always
# visible without scrolling.
progress = status_placeholder.progress(0, text="Loading model...")
# Lazy imports
from betaearth import BetaEarth
import torch
from betaearth.generate import (
compute_grid, download_dem, download_s2, download_s1,
download_s2_cloud_mask,
_search_stac, _seasonal_select, check_coverage,
write_geotiff, fit_pca, write_pca_preview,
)
device = "cuda" if torch.cuda.is_available() else "cpu"
model = BetaEarth.from_pretrained(device=device)
progress.progress(5, text=f"Model loaded on {device}")
user_grid = compute_grid(bbox)
was_padded, pad_px_x, pad_px_y = expand_to_min(bbox)
grid = pad_grid_from_user(user_grid, pad_px_x, pad_px_y) if was_padded else user_grid
h, w = grid["shape"]
out_h, out_w = user_grid["shape"]
run_id = uuid.uuid4().hex[:8]
output_dir = Path(tempfile.mkdtemp()) / f"betaearth_{run_id}"
output_dir.mkdir(parents=True, exist_ok=True)
# DEM (shared across years) — download at padded size, save cropped
progress.progress(8, text="Downloading DEM...")
dem = download_dem(grid)
if save_per_timestamp_input:
dem_out = crop_to_user(dem, grid, user_grid, channel_axis=0)
write_geotiff(dem_out.astype(np.float32), user_grid, output_dir / "dem.tif", band_first=True)
# DEM preview (from cropped array)
from PIL import Image
d = dem_out[0].astype(np.float32)
lo, hi = np.percentile(d[np.isfinite(d)], [2, 98])
d_norm = np.clip((d - lo) / max(hi - lo, 1e-6), 0, 1)
Image.fromarray((d_norm * 255).astype(np.uint8)).save(output_dir / "dem_preview.png")
all_previews = []
all_summaries = []
last_annual_preview = None
for yi, year in enumerate(years):
# Progress range for this year: spread evenly across [10, 90]
yr_lo = 10 + int(80 * yi / n_years)
yr_hi = 10 + int(80 * (yi + 1) / n_years)
yr_label = f"[{year}] " if n_years > 1 else ""
files_dir = output_dir / f"{year}_files"
files_dir.mkdir()
# Search. We deliberately don't filter by cloud at search time — the
# filter is applied per-quarter in _seasonal_select with a fallback
# so summer quarters don't get dropped entirely in cloudy regions.
progress.progress(yr_lo, text=f"{yr_label}Searching Planetary Computer...")
if custom_dates:
import pystac_client, planetary_computer
from datetime import date as _date
cd_start = _date.fromisoformat(custom_dates[0])
cd_end = _date.fromisoformat(custom_dates[1])
yr_start = max(cd_start, _date(year, 1, 1))
yr_end = min(cd_end, _date(year, 12, 31))
catalog = pystac_client.Client.open(
"https://planetarycomputer.microsoft.com/api/stac/v1",
modifier=planetary_computer.sign_inplace,
)
s2_items = list(catalog.search(
collections=["sentinel-2-l2a"], bbox=list(bbox),
datetime=f"{yr_start}/{yr_end}",
max_items=200,
).items())
s1_items = list(catalog.search(
collections=["sentinel-1-rtc"], bbox=list(bbox),
datetime=f"{yr_start}/{yr_end}",
max_items=200,
).items())
else:
s2_items = _search_stac(bbox, year, "sentinel-2-l2a")
s1_items = _search_stac(bbox, year, "sentinel-1-rtc")
s2_items = _seasonal_select(s2_items, max_per_quarter=6, max_cloud=max_cloud)
s1_items = _seasonal_select(s1_items, max_per_quarter=1, use_cloud=False)
total_found = len(s2_items) + len(s1_items)
# Process scenes
emb_sum = np.zeros((h, w, 64), dtype=np.float64)
emb_count = np.zeros((h, w), dtype=np.int32)
used_scenes = []
all_items = [(it, "S2") for it in s2_items] + [(it, "S1") for it in s1_items]
n_total = len(all_items)
skipped = 0
failed = 0
for i, (item, sensor) in enumerate(all_items):
pct = yr_lo + int((yr_hi - yr_lo - 10) * i / max(n_total, 1))
dt = item.datetime
doy = dt.timetuple().tm_yday
try:
if sensor == "S2":
mgrs = item.properties.get("s2:mgrs_tile", "???")
cc = item.properties.get("eo:cloud_cover", 0)
progress.progress(pct, text=f"{yr_label}[{i+1}/{n_total}] S2 {mgrs} {dt.date()} (cloud={cc:.0f}%)")
data = download_s2(item, grid)
else:
progress.progress(pct, text=f"{yr_label}[{i+1}/{n_total}] S1 {dt.date()}")
data = download_s1(item, grid)
cov = check_coverage(data)
if cov < min_coverage:
skipped += 1
continue
progress.progress(pct + 1, text=f"{yr_label}[{i+1}/{n_total}] Predicting {sensor} {dt.date()}...")
if sensor == "S2":
emb = model.predict(s2_l2a=data, dem=dem, doy=doy, tile_size=224, overlap=112)
ts_label = f"{dt.date()}_s2"
# SCL-based per-pixel cloud/shadow mask (None if asset missing)
cloud_mask = download_s2_cloud_mask(item, grid)
else:
emb = model.predict(s1=data, dem=dem, doy=doy, tile_size=224, overlap=112)
ts_label = f"{dt.date()}_s1"
cloud_mask = None
valid = np.linalg.norm(emb, axis=-1) > 1e-6
if cloud_mask is not None:
valid &= cloud_mask
emb_sum[valid] += emb[valid]
emb_count[valid] += 1
except Exception as _scene_err: # noqa: BLE001
failed += 1
print(
f"[scene] FAILED {sensor} {getattr(item, 'id', '?')}: "
f"{type(_scene_err).__name__}: {_scene_err}",
flush=True,
)
continue
ts_dir = files_dir / ts_label
ts_dir.mkdir(parents=True, exist_ok=True)
if save_per_timestamp:
emb_out = crop_to_user(emb, grid, user_grid, channel_axis=-1)
write_geotiff(emb_out.astype(np.float32), user_grid, ts_dir / "embedding.tif")
if save_per_timestamp_input:
# data is band-first: (C, H, W) — crop band-first
data_out = crop_to_user(data, grid, user_grid, channel_axis=0)
write_geotiff(data_out.astype(np.float32), user_grid, ts_dir / "input.tif", band_first=True)
# RGB preview (from cropped input)
from PIL import Image
if sensor == "S2":
# S2 band order: [B02, B03, B04, B08, B05, B06, B07, B11, B12]
# RGB = B04, B03, B02 → indices 2, 1, 0
rgb = np.stack([data_out[2], data_out[1], data_out[0]], axis=-1).astype(np.float32)
rgb = np.clip(rgb / 3000.0, 0, 1) ** (1/2.2)
else:
# S1: VV, VH → composite with ratio for third channel
vv = 10 * np.log10(np.clip(data_out[0], 1e-6, None))
vh = 10 * np.log10(np.clip(data_out[1], 1e-6, None))
ratio = vv - vh
def _norm(x):
lo, hi = np.percentile(x[np.isfinite(x)], [2, 98])
return np.clip((x - lo) / max(hi - lo, 1e-6), 0, 1)
rgb = np.stack([_norm(vv), _norm(vh), _norm(ratio)], axis=-1)
Image.fromarray((rgb * 255).astype(np.uint8)).save(ts_dir / "preview_rgb.png")
used_scenes.append({"sensor": sensor, "date": str(dt.date()), "doy": doy, "coverage": round(cov, 1)})
if not used_scenes:
all_summaries.append(f"**{year}:** No scenes passed {min_coverage}% filter")
continue
# Average
progress.progress(yr_hi - 5, text=f"{yr_label}Averaging...")
covered = emb_count > 0
avg = np.zeros((h, w, 64), dtype=np.float32)
avg[covered] = (emb_sum[covered] / emb_count[covered, np.newaxis]).astype(np.float32)
norms = np.linalg.norm(avg, axis=-1, keepdims=True)
avg = avg / np.clip(norms, 1e-8, None)
avg[~covered] = 0
avg_out = crop_to_user(avg, grid, user_grid, channel_axis=-1)
write_geotiff(avg_out, user_grid, output_dir / f"{year}.tif")
# PCA previews (fit on cropped annual, apply same basis to timestamps)
progress.progress(yr_hi - 2, text=f"{yr_label}PCA previews...")
pca_state = fit_pca(avg_out)
annual_preview = output_dir / f"{year}_preview_pca.png"
write_pca_preview(avg_out, annual_preview, pca_state=pca_state)
last_annual_preview = str(annual_preview)
import rasterio
for ts_dir in sorted(files_dir.iterdir()):
emb_tif = ts_dir / "embedding.tif"
if emb_tif.exists():
with rasterio.open(emb_tif) as src:
ts_emb = src.read().transpose(1, 2, 0)
write_pca_preview(ts_emb, ts_dir / "preview_pca.png", pca_state=pca_state)
# Manifest
manifest = {
"bbox_4326": list(bbox), "year": year,
"min_coverage": min_coverage,
"n_scenes_found": total_found,
"n_scenes_used": len(used_scenes),
"scenes": used_scenes,
}
with open(output_dir / f"{year}_manifest.json", "w") as f:
json.dump(manifest, f, indent=2)
all_previews.append((str(annual_preview), f"{year} average"))
for ts_dir in sorted(files_dir.iterdir()):
png = ts_dir / "preview_pca.png"
if png.exists():
all_previews.append((str(png), f"{year}/{ts_dir.name}"))
rgb_png = ts_dir / "preview_rgb.png"
if rgb_png.exists():
all_previews.append((str(rgb_png), f"{year}/{ts_dir.name} (RGB input)"))
all_summaries.append(
f"**{year}:** {len(used_scenes)} scenes "
f"({total_found} found, {skipped} skipped)"
)
# ZIP everything
progress.progress(92, text="Creating download archive...")
zip_path = output_dir.parent / f"betaearth_{run_id}"
shutil.make_archive(str(zip_path), "zip", str(output_dir))
zip_file = str(zip_path) + ".zip"
with open(zip_file, "rb") as f:
zip_data = f.read()
# Read previews into memory so we can clean up the tempdir immediately.
# Each PNG is ~4 MB — cheaper than leaving a multi-GB tempdir per user.
previews_in_mem = []
annual_preview_bytes = None
for ppath, plabel in all_previews:
p = Path(ppath)
if p.exists():
previews_in_mem.append((p.read_bytes(), plabel))
if last_annual_preview and Path(last_annual_preview).exists():
annual_preview_bytes = Path(last_annual_preview).read_bytes()
st.session_state.results = {
"bbox": list(bbox),
"zip_data": zip_data,
"zip_name": f"betaearth_{run_id}.zip",
"annual_preview": annual_preview_bytes,
"summary": "\n\n".join(all_summaries),
"previews": previews_in_mem,
}
# Clean up tempdir now that everything we need is in session_state.
# Previously these accumulated across every Generate click until the
# Space was restarted (bug 5).
try:
shutil.rmtree(output_dir.parent, ignore_errors=True)
except Exception:
pass
progress.progress(100, text="Done!")
st.rerun()
# ---------------------------------------------------------------------------
# Overlay controls + preview gallery (below map, only after generation)
# ---------------------------------------------------------------------------
if "results" in st.session_state and st.session_state.results:
res = st.session_state.results
previews = res.get("previews", [])
if previews:
labels = [label for _, label in previews]
col_slider, col_opacity, col_dl = st.columns([3, 1, 1])
with col_slider:
st.select_slider("Preview frame", options=labels, value=labels[0], key="preview_frame")
with col_opacity:
st.slider("Opacity", 0.0, 1.0, 0.7, step=0.05, key="opacity")
with col_dl:
st.write("") # vertical spacing to align with sliders
st.download_button(
"📦 Download ZIP",
data=res["zip_data"],
file_name=res["zip_name"],
mime="application/zip",
use_container_width=True,
)
st.subheader("PCA-RGB Previews")
cols = st.columns(min(len(previews), 5))
for i, (img_bytes, label) in enumerate(previews):
with cols[i % len(cols)]:
st.image(img_bytes, caption=label, use_container_width=True)