Download quickstart.py from Jwjohnson314/Aurora-FM: direct link, hf CLI and curl.
- Browser
- Download file 23.9 kB
-
https://huggingface.co/Jwjohnson314/Aurora-FM/resolve/main/quickstart.py
- Command line
-
hf download hf://Jwjohnson314/Aurora-FM/quickstart.py
-
curl -L -o quickstart.py https://huggingface.co/Jwjohnson314/Aurora-FM/resolve/main/quickstart.py
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() | |