Aurora-FM / quickstart.py
Jwjohnson314's picture
Upload 2 files
7eed357 verified
Raw History Blame Contribute Delete
23.9 kB
"""
quickstart.py
=============
A small, **self-contained** entry point for classifying a single THEMIS
All-Sky Imager (ASI) frame with the fine-tuned aurora foundation model.
What this module does
---------------------
Given a *site*, a *datetime*, and a *frame index* it will:
1. **Fetch** the corresponding THEMIS level-1 all-sky CDF file. If you
already have the data locally, point ``data_dir`` (or ``--data-dir``) at
it and no download happens. Otherwise the file is downloaded once from
the Berkeley THEMIS data server and cached for next time.
2. **Load** the requested video frame out of that CDF.
3. **Preprocess** the raw frame exactly the way the model's training /
inference pipeline does (Clausen percentile normalization, circular
field-of-view mask, green-channel packing, resize to 224x224).
4. **Classify** the frame with the fine-tuned SimCLR model, returning a
probability for each of the six auroral classes.
The model
---------
The classifier is the fine-tuned checkpoint
``aurora-fm-finetuned.tar``
which is the fine-tuned counterpart of the aurora foundation model
``aurora-fm.tar``. ``aurora-fm.tar`` is a
*self-supervised* SimCLR model: a ResNet-18 encoder plus a contrastive
projection head (``projector.0`` / ``projector.2``). On its own it only
produces 512-d feature vectors -- it does not classify. The
``...finetuned.tar`` checkpoint takes that same encoder and fine-tunes
it end-to-end with a two-layer classification head
(``projector1``: 512->512, ``projector2``: 512->6) on the Oslo Aurora THEMIS (OATH) dataset (Clausen et al. 2018), so a single forward
pass followed by a softmax yields the six class probabilities directly.
The six classes (fixed order, matching the training labels) are:
arc, diffuse, discrete, cloudy, moon, clear
Command-line usage
------------------
python quickstart.py --site fsmi --datetime 20140326T09 --frame 0
Programmatic usage
------------------
from quickstart import load_model, classify
model = load_model() # load the classifier once
result = classify("fsmi", "20140326T09", 0, model=model)
print(result["label"], result["probabilities"])
"""
from __future__ import annotations
import argparse
import os
from pathlib import Path
from typing import Optional
import cdflib
import numpy as np
import torch
import torch.nn as nn
import torchvision
import requests
from PIL import Image
# --------------------------------------------------------------------------- #
# Constants
# --------------------------------------------------------------------------- #
ROOT = Path(__file__).parent.resolve()
# Default fine-tuned classifier checkpoint. Paths are resolved relative to
# this file so the tool works from any working directory and on any machine.
# Override at runtime with the ``--checkpoint`` flag, the ``checkpoint_path``
# argument to ``load_model``, or the ``THEMIS_AURORA_CHECKPOINT`` env var.
DEFAULT_CHECKPOINT = Path(
os.environ.get(
"THEMIS_AURORA_CHECKPOINT",
ROOT / "important-checkpoints" / "aurora-fm-finetuned.tar",
)
)
# Where downloaded CDFs are cached. Override with the ``cache_dir`` argument
# to ``fetch_cdf`` / ``classify``, the ``--cache-dir`` flag, or the
# ``THEMIS_AURORA_CACHE_DIR`` env var.
DEFAULT_CACHE_DIR = Path(
os.environ.get("THEMIS_AURORA_CACHE_DIR", ROOT / "cdf_cache")
)
# Optional directory of pre-existing local CDFs (so users who already have the
# data need not re-download it). Override with the ``data_dir`` argument to
# ``fetch_cdf`` / ``classify``, the ``--data-dir`` flag, or the
# ``THEMIS_AURORA_DATA_DIR`` env var. ``None`` means "no local data dir".
_env_data_dir = os.environ.get("THEMIS_AURORA_DATA_DIR")
DEFAULT_DATA_DIR = Path(_env_data_dir) if _env_data_dir else None
# Base URL for THEMIS level-1 all-sky "full" (asf) imagery.
VIDEO_BASE = "http://themis.ssl.berkeley.edu/data/themis/thg/l1/asi/"
# The six auroral classes, in the exact order the model's output units were
# trained on (see inference() in large_scale_inference.py).
CLASS_NAMES = ["arc", "diffuse", "discrete", "cloudy", "moon", "clear"]
# The 24 THEMIS ASI ground stations (from site_list.py).
SITES = [
"atha", "chbg", "ekat", "fsim", "fsmi", "fykn", "gako", "gbay",
"gill", "inuv", "kapu", "kian", "kuuj", "mcgr", "nrsq", "pgeo",
"pina", "rank", "snap", "snkq", "talo", "tpas", "whit", "yknf",
]
# Preprocessing / model input constants (from simclr_config.yaml and the
# inference default in large_scale_inference.py).
IMAGE_SIZE = 224 # network input side length in pixels
RADIUS = 0.85 # circular field-of-view crop, as a fraction of the half-width
RAW_SIZE = 256 # native THEMIS ASI frame side length in pixels
# Network timeout for HTTP HEAD / GET requests, in seconds.
TIMEOUT = 600
# --------------------------------------------------------------------------- #
# Datetime helpers
# --------------------------------------------------------------------------- #
def _parse_datetime(datetime_str: str) -> dict:
"""
Parse a ``YYYYMMDDTHH`` datetime string into its components.
THEMIS all-sky CDFs are stored one file per site *per hour*, so the finest
granularity we ever need is the hour.
:param datetime_str: e.g. ``"20140326T09"`` (year 2014, month 03, day 26,
hour 09 UTC).
:return: dict with integer ``year``, ``month``, ``day``, ``hour`` and a
``stamp`` string ``"YYYYMMDDHH"`` used to build filenames / URLs.
:raises ValueError: if the string is not in the expected format.
"""
s = datetime_str.strip().upper()
if "T" not in s:
raise ValueError(
f"datetime {datetime_str!r} must be in 'YYYYMMDDTHH' format, "
f"e.g. '20140326T09'"
)
date_part, hour_part = s.split("T")
if len(date_part) != 8 or len(hour_part) != 2:
raise ValueError(
f"datetime {datetime_str!r} must be in 'YYYYMMDDTHH' format, "
f"e.g. '20140326T09' (8-digit date, 'T', 2-digit hour)"
)
try:
year = int(date_part[0:4])
month = int(date_part[4:6])
day = int(date_part[6:8])
hour = int(hour_part)
except ValueError as exc:
raise ValueError(f"could not parse datetime {datetime_str!r}: {exc}") from exc
return {
"year": year,
"month": month,
"day": day,
"hour": hour,
# Compact YYYYMMDDHH stamp used in the CDF filename and remote URL.
"stamp": f"{year:04d}{month:02d}{day:02d}{hour:02d}",
}
def _cdf_filename(site: str, datetime_str: str) -> str:
"""Return the canonical CDF filename for a site/hour, e.g.
``thg_l1_asf_fsmi_2014032609_v01.cdf``."""
stamp = _parse_datetime(datetime_str)["stamp"]
return f"thg_l1_asf_{site}_{stamp}_v01.cdf"
def _cdf_url(site: str, datetime_str: str) -> str:
"""Build the remote URL for a site/hour CDF on the THEMIS data server.
Mirrors ``util.urlgen`` / ``util.fetch_cdf``:
``{VIDEO_BASE}{site}/{YYYY}/{MM}/thg_l1_asf_{site}_{YYYYMMDDHH}_v01.cdf``.
"""
p = _parse_datetime(datetime_str)
return (
f"{VIDEO_BASE}{site}/{p['year']:04d}/{p['month']:02d}/"
f"{_cdf_filename(site, datetime_str)}"
)
# --------------------------------------------------------------------------- #
# Data fetching / loading
# --------------------------------------------------------------------------- #
def _find_local_cdf(site: str, datetime_str: str, search_dir: Path) -> Optional[Path]:
"""
Look for an already-downloaded CDF under ``search_dir``.
Two common on-disk layouts are supported:
* **flat** -- all CDFs directly in ``search_dir``:
``search_dir/thg_l1_asf_fsmi_2014032609_v01.cdf``
* **Berkeley-mirrored** -- the same ``{site}/{YYYY}/{MM}/`` sub-tree the
THEMIS server uses:
``search_dir/fsmi/2014/03/thg_l1_asf_fsmi_2014032609_v01.cdf``
:return: the path if found, else ``None``.
"""
fname = _cdf_filename(site, datetime_str)
p = _parse_datetime(datetime_str)
candidates = [
Path(search_dir) / fname,
Path(search_dir) / site / f"{p['year']:04d}" / f"{p['month']:02d}" / fname,
]
for candidate in candidates:
if candidate.exists():
return candidate
return None
def fetch_cdf(
site: str,
datetime_str: str,
data_dir: Optional[Path] = DEFAULT_DATA_DIR,
cache_dir: Path = DEFAULT_CACHE_DIR,
download: bool = True,
) -> Path:
"""
Return a local path to the requested THEMIS ASI CDF.
Resolution order:
1. If ``data_dir`` is given and already contains the file (flat or
Berkeley-mirrored layout), use it as-is -- nothing is copied or
downloaded.
2. Otherwise, if the file is already in ``cache_dir``, use that.
3. Otherwise, if ``download`` is ``True``, download it from the Berkeley
THEMIS server into ``cache_dir``.
:param site: 4-character site code, e.g. ``"fsmi"``.
:param datetime_str: ``YYYYMMDDTHH`` UTC hour, e.g. ``"20140326T09"``.
:param data_dir: optional directory of pre-existing local CDFs to search
before downloading. ``None`` disables the local lookup.
:param cache_dir: directory used to store downloaded CDFs.
:param download: if ``False``, never hit the network -- raise instead of
downloading a missing file.
:return: path to the local CDF file.
:raises RuntimeError: if the file is not found locally and either the
download is disabled or fails (e.g. no data exists for that
site/hour -> HTTP 404).
"""
site = site.lower()
# 1. Prefer an existing local copy from the user's data directory.
if data_dir is not None:
found = _find_local_cdf(site, datetime_str, Path(data_dir))
if found is not None:
print(f"Using local CDF: {found}")
return found
cache_dir = Path(cache_dir)
# 2. Fall back to the download cache.
local_path = cache_dir / _cdf_filename(site, datetime_str)
if local_path.exists():
print(f"Using cached CDF: {local_path}")
return local_path
# 3. Download, unless the caller opted out of network access.
if not download:
raise RuntimeError(
f"CDF for site {site!r} at {datetime_str!r} not found locally "
f"(searched data_dir={data_dir} and cache_dir={cache_dir}) and "
f"download=False."
)
cache_dir.mkdir(parents=True, exist_ok=True)
url = _cdf_url(site, datetime_str)
print(f"Downloading {url} ...")
response = requests.get(url, allow_redirects=True, timeout=TIMEOUT)
if response.status_code != 200:
raise RuntimeError(
f"could not download {url}: HTTP {response.status_code}. "
f"There is probably no data for site {site!r} at {datetime_str!r}."
)
# Write atomically-ish: download to a temp name, then rename, so an
# interrupted download never leaves a truncated file in the cache.
tmp_path = local_path.with_suffix(local_path.suffix + ".part")
tmp_path.write_bytes(response.content)
tmp_path.rename(local_path)
print(f"Download complete: {local_path}")
return local_path
def load_themis_cdf(cdf_path: Path):
"""
Read a THEMIS all-sky CDF and return its image stack and timestamps.
The variable names inside the CDF are keyed by site, e.g. ``thg_asf_fsmi``
(the image cube, shape ``[n_frames, 256, 256]``) and ``thg_asf_fsmi_epoch``
(the per-frame timestamps).
:param cdf_path: path to a local ``thg_l1_asf_*.cdf`` file.
:return: tuple ``(imgs, times)`` where ``imgs`` is a
``[n_frames, 256, 256]`` numpy array and ``times`` is a numpy array of
``datetime64`` timestamps of length ``n_frames``.
:raises ValueError: if the CDF contains no frames.
"""
cdf_path = Path(cdf_path)
# The site code is the 4th underscore-delimited token of the filename:
# thg_l1_asf_<site>_<stamp>_v01.cdf
site = cdf_path.name.split("_")[3]
handle = cdflib.cdfread.CDF(cdf_path)
times = cdflib.cdfepoch.to_datetime(handle[f"thg_asf_{site}_epoch"][:])
imgs = handle[f"thg_asf_{site}"][:]
# A single-frame file loads as a 2-D array; promote it to a 1-frame stack.
if imgs.ndim == 2:
imgs = imgs[None, ...]
if len(times) == 0:
raise ValueError(f"no frames found in {cdf_path}")
return imgs, np.asarray(times)
# --------------------------------------------------------------------------- #
# Preprocessing
# --------------------------------------------------------------------------- #
def _circular_mask(radius: float = RADIUS, size: int = RAW_SIZE) -> np.ndarray:
"""
Build a boolean circular field-of-view mask.
THEMIS all-sky cameras image the full sky through a fisheye lens, so only a
central disc contains sky; the corners are housing / horizon. We keep a
disc of radius ``radius * (size / 2)`` pixels centered on the frame and
zero everything outside it. ``radius = 0.85`` was chosen to roughly match
the field of view of the OATH training images.
:return: ``[size, size]`` boolean array, ``True`` inside the disc.
"""
dist = int(radius * (size // 2))
yy, xx = np.ogrid[:size, :size]
dist_from_center = np.sqrt((xx - size // 2) ** 2 + (yy - size // 2) ** 2)
return dist_from_center <= dist
# Precompute the mask and the test-time image transform once at import.
_MASK = _circular_mask()
# Test-time transform: resize the 256x256 packed image to the 224x224 network
# input and convert to a CxHxW float tensor in [0, 1]. This is exactly
# ``model.TransformsSimCLR(size=IMAGE_SIZE).test_transform``.
_TEST_TRANSFORM = torchvision.transforms.Compose(
[
torchvision.transforms.Resize(size=IMAGE_SIZE),
torchvision.transforms.ToTensor(),
]
)
def preprocess_frame(raw_frame: np.ndarray) -> torch.Tensor:
"""
Turn one raw 256x256 THEMIS frame into a model-ready input tensor.
Step by step:
1. **Clausen percentile normalization** -- subtract the 1st percentile,
divide by the 99th percentile, and clip to ``[0, 1]``. This gives a
robust contrast stretch that ignores hot pixels and the dark floor
(Clausen & Nickisch, 2018).
2. **Circular mask** -- zero everything outside the fisheye field of view.
3. **Green-channel packing** -- place the single-channel image in the green
channel of an otherwise-black RGB image.
4. **To PIL uint8**, then apply the test transform (resize to 224, ToTensor).
:param raw_frame: ``[256, 256]`` array of raw sensor counts (any integer or
float dtype).
:return: a ``[3, 224, 224]`` float tensor in ``[0, 1]``, ready to be
batched and fed to the model.
"""
imarray = np.asarray(raw_frame).astype(np.float32)
# 1. Clausen percentile normalization.
p1 = np.percentile(imarray, 1)
imarray = imarray - p1
p99 = np.percentile(imarray, 99)
# Guard against a degenerate (flat) frame producing a divide-by-zero.
if p99 != 0:
imarray = imarray / p99
imarray = np.clip(imarray, 0, 1)
# 2. Circular field-of-view mask.
imarray[~_MASK] = 0
# 3. Pack into the green channel of an RGB image.
rgb = np.zeros((RAW_SIZE, RAW_SIZE, 3), dtype=np.float32)
rgb[:, :, 1] = imarray
# 4. To PIL uint8, then resize + ToTensor.
pil_img = Image.fromarray((rgb * 255).astype(np.uint8))
return _TEST_TRANSFORM(pil_img)
# --------------------------------------------------------------------------- #
# Model
# --------------------------------------------------------------------------- #
class _Identity(nn.Module):
"""Pass-through module used to strip the ResNet's classification head so
the encoder emits its 512-d pooled features."""
def forward(self, x):
return x
class Finetuner(nn.Module):
"""
The fine-tuned aurora classifier.
A ResNet-18 encoder (its final fc replaced by identity, so it outputs a
512-d feature vector) followed by a two-layer classification head:
projector1: Linear(512 -> 512, no bias) -> ReLU
projector2: Linear(512 -> n_classes, no bias)
A forward pass returns raw class logits; apply softmax for probabilities.
"""
def __init__(self, encoder: nn.Module, n_classes: int, n_features: int):
super().__init__()
self.encoder = encoder
self.n_features = n_features
self.encoder.fc = _Identity()
self.projector1 = nn.Linear(self.n_features, self.n_features, bias=False)
self.projector2 = nn.Linear(self.n_features, n_classes, bias=False)
def forward(self, x):
x = nn.functional.relu(self.projector1(self.encoder(x)))
return self.projector2(x)
def load_model(
checkpoint_path: Path = DEFAULT_CHECKPOINT,
device: Optional[torch.device] = None,
) -> nn.Module:
"""
Build the ResNet-18 ``Finetuner`` and load the fine-tuned weights.
:param checkpoint_path: path to the ``...-finetuned.tar`` state dict.
Defaults to the checkpoint-60 fine-tuned model.
:param device: torch device to place the model on. Defaults to CUDA if
available, else CPU.
:return: the model in ``eval`` mode, on ``device``.
:raises FileNotFoundError: if the checkpoint does not exist.
"""
checkpoint_path = Path(checkpoint_path)
if not checkpoint_path.exists():
raise FileNotFoundError(f"checkpoint not found: {checkpoint_path}")
if device is None:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Build a headless ResNet-18 encoder; 512 = encoder.fc.in_features.
encoder = torchvision.models.resnet18(weights=None)
n_features = encoder.fc.in_features
model = Finetuner(encoder, n_classes=len(CLASS_NAMES), n_features=n_features)
print(f"Loading classifier weights from {checkpoint_path} ...")
state_dict = torch.load(checkpoint_path, map_location=device)
model.load_state_dict(state_dict)
model = model.to(device)
model.eval()
return model
# --------------------------------------------------------------------------- #
# End-to-end classification
# --------------------------------------------------------------------------- #
def classify(
site: str,
datetime_str: str,
frame: int,
model: Optional[nn.Module] = None,
data_dir: Optional[Path] = DEFAULT_DATA_DIR,
cache_dir: Path = DEFAULT_CACHE_DIR,
download: bool = True,
return_frame: bool = False,
) -> dict:
"""
Fetch, preprocess and classify a single THEMIS ASI frame.
:param site: 4-character site code, e.g. ``"fsmi"`` (case-insensitive).
:param datetime_str: UTC hour in ``YYYYMMDDTHH`` format, e.g.
``"20140326T09"``.
:param frame: index of the frame within that hour's CDF (0-based). A
THEMIS ASI hour typically holds ~1200 frames at 3-second cadence.
:param model: a preloaded model. If ``None`` a
model is loaded on each call -- pass one in when classifying many
frames so the weights are only loaded once.
:param data_dir: optional directory of pre-existing local CDFs to search
before downloading.
:param cache_dir: directory for cached CDF downloads.
:param download: if ``False``, never download -- require the CDF to be
present in ``data_dir`` or ``cache_dir``.
:param return_frame: if ``True``, also return the raw 256x256 frame under
the ``"raw_frame"`` key (handy for plotting alongside the prediction).
:return: dict with keys:
``site``, ``datetime``, ``frame``, ``time`` (the frame's UTC
timestamp), ``label`` (top-1 class name), ``probabilities`` (dict
mapping each class name to its probability), and optionally
``raw_frame``.
:raises IndexError: if ``frame`` is out of range for the loaded CDF.
"""
site = site.lower()
if model is None:
model = load_model()
device = next(model.parameters()).device
# 1. Locate the CDF (local data dir, cache, or download) and read it.
cdf_path = fetch_cdf(
site, datetime_str, data_dir=data_dir, cache_dir=cache_dir, download=download
)
imgs, times = load_themis_cdf(cdf_path)
n_frames = imgs.shape[0]
if not (0 <= frame < n_frames):
raise IndexError(
f"frame {frame} out of range: {site} {datetime_str} has "
f"{n_frames} frames (valid indices 0..{n_frames - 1})"
)
# 2. Grab the requested frame and preprocess it.
raw_frame = np.asarray(imgs[frame])
tensor = preprocess_frame(raw_frame)
# 3. Classify: add a batch dimension, forward, softmax.
batch = tensor.unsqueeze(0).to(device)
with torch.no_grad():
logits = model(batch)
probs = torch.softmax(logits, dim=1).squeeze(0).cpu().numpy()
probabilities = {name: float(p) for name, p in zip(CLASS_NAMES, probs)}
label = CLASS_NAMES[int(np.argmax(probs))]
result = {
"site": site,
"datetime": datetime_str,
"frame": frame,
"time": times[frame],
"label": label,
"probabilities": probabilities,
}
if return_frame:
result["raw_frame"] = raw_frame
return result
# --------------------------------------------------------------------------- #
# Command-line interface
# --------------------------------------------------------------------------- #
def _build_arg_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description=(
"Classify a single THEMIS all-sky imager frame with the "
"fine-tuned aurora model."
)
)
parser.add_argument(
"--site",
required=True,
help=f"4-character THEMIS site code. One of: {', '.join(SITES)}",
)
parser.add_argument(
"--datetime",
required=True,
help="UTC hour in YYYYMMDDTHH format, e.g. 20140326T09",
)
parser.add_argument(
"--frame",
type=int,
default=0,
help="0-based frame index within the hour's CDF (default: 0)",
)
parser.add_argument(
"--checkpoint",
default=str(DEFAULT_CHECKPOINT),
help="path to the fine-tuned classifier checkpoint (.tar)",
)
parser.add_argument(
"--data-dir",
default=(str(DEFAULT_DATA_DIR) if DEFAULT_DATA_DIR else None),
help="directory of pre-existing local CDFs to use instead of "
"downloading (flat or {site}/{YYYY}/{MM}/ layout)",
)
parser.add_argument(
"--cache-dir",
default=str(DEFAULT_CACHE_DIR),
help="directory for cached CDF downloads",
)
parser.add_argument(
"--no-download",
action="store_true",
help="never download; require the CDF to be present locally",
)
return parser
def main() -> None:
args = _build_arg_parser().parse_args()
if args.site.lower() not in SITES:
print(f"warning: {args.site!r} is not a known THEMIS site; continuing anyway")
model = load_model(checkpoint_path=Path(args.checkpoint))
result = classify(
site=args.site,
datetime_str=args.datetime,
frame=args.frame,
model=model,
data_dir=Path(args.data_dir) if args.data_dir else None,
cache_dir=Path(args.cache_dir),
download=not args.no_download,
)
print()
print(f"Site : {result['site']}")
print(f"Datetime : {result['datetime']} (frame {result['frame']})")
print(f"Timestamp : {result['time']}")
print(f"Prediction: {result['label']}")
print("Probabilities:")
for name, prob in sorted(
result["probabilities"].items(), key=lambda kv: kv[1], reverse=True
):
print(f" {name:9s} {prob:6.3f}")
if __name__ == "__main__":
main()