Download orbit2vec.py from jgalego/orbit2vec: direct link, hf CLI and curl.
- Browser
- Download file 34.6 kB
-
https://huggingface.co/jgalego/orbit2vec/resolve/main/orbit2vec.py
- Command line
-
hf download hf://jgalego/orbit2vec/orbit2vec.py
-
curl -L -o orbit2vec.py https://huggingface.co/jgalego/orbit2vec/resolve/main/orbit2vec.py
34.6 kB
| # /// script | |
| # requires-python = ">=3.12" | |
| # dependencies = [ | |
| # "datasets", | |
| # "huggingface-hub", | |
| # "jinja2", | |
| # "matplotlib", | |
| # "numpy", | |
| # "polars", | |
| # "safetensors", | |
| # "sgp4", | |
| # "torch", | |
| # "umap-learn", | |
| # ] | |
| # /// | |
| """Orbit2Vec: embed satellite orbit histories so that windows of the same object land together.""" | |
| # pylint: disable=too-many-arguments,too-many-instance-attributes,too-many-locals | |
| # pylint: disable=too-many-positional-arguments,too-many-statements | |
| import argparse | |
| import json | |
| import math | |
| import os | |
| import random | |
| import sys | |
| import time | |
| from collections import Counter | |
| from pathlib import Path | |
| def cpus(): | |
| """Return the CPU quota visible to this process.""" | |
| try: | |
| quota, period = Path("/sys/fs/cgroup/cpu.max").read_text(encoding="utf-8").split() | |
| if quota != "max": | |
| return max(1, int(quota) // int(period)) | |
| except OSError: | |
| pass | |
| return len(os.sched_getaffinity(0)) | |
| # Polars sizes its thread pool on import from the host's cores, not the job's quota. | |
| os.environ.setdefault("POLARS_MAX_THREADS", str(cpus())) | |
| # pylint: disable=wrong-import-position | |
| import datasets | |
| import huggingface_hub.utils | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import polars as pl | |
| import torch | |
| import umap | |
| from datasets import Dataset | |
| from huggingface_hub import ( | |
| HfApi, | |
| ModelCard, | |
| ModelCardData, | |
| PyTorchModelHubMixin, | |
| hf_hub_download, | |
| snapshot_download, | |
| ) | |
| from huggingface_hub.errors import RepositoryNotFoundError | |
| from sgp4.api import Satrec | |
| from torch import nn | |
| from torch.nn import functional | |
| REPO = "jgalego/orbit2vec" | |
| DATA = "jgalego/orbit2vec-windows" | |
| TLE = "juliensimon/space-track-tle-history" | |
| SATCAT = "juliensimon/space-track-satcat" | |
| MANEUVERS = "RhynoWu/starlink-maneuver-db" | |
| HERE = Path(__file__).parent | |
| # Progress bars redraw in place, which shows up as garbage in HF Jobs logs. | |
| PROGRESS = sys.stderr.isatty() | |
| ELEMENTS = [ | |
| "mean_motion", | |
| "eccentricity", | |
| "inclination", | |
| "raan", | |
| "arg_perigee", | |
| "bstar", | |
| "mean_motion_dot", | |
| ] | |
| FEATURES = 11 | |
| REGIMES = ["LEO", "MEO", "GEO", "HEO"] | |
| MU, RE = 398600.4418, 6378.137 | |
| # Windows are 32 UTC days on a grid starting 1970-01-01. Training pairs end before window 639 | |
| # (2025-12-26); the test pair is windows 641 and 642 (2026-02-28 to 2026-05-02). | |
| WINDOW = 32 | |
| TRAIN_END = 639 | |
| TEST = 641 | |
| MIN_DAYS = 8 | |
| PROBE_DAYS = 10 | |
| NON_PAYLOAD = {"DEB": "debris", "R/B": "rocket body", "UNK": "unknown"} | |
| def semi_major(n): | |
| """Semi-major axis in km from mean motion in revolutions per day.""" | |
| return (MU / (n * 2 * math.pi / 86400) ** 2) ** (1 / 3) | |
| def features(x): | |
| """Per-day inputs from raw mean elements (NaN on days without a TLE), and the padding mask.""" | |
| n, e, i, raan, argp, bstar, ndot = x.unbind(-1) | |
| a = semi_major(n) | |
| da = a - a.nanmedian(1, keepdim=True).values | |
| di = i - i.nanmedian(1, keepdim=True).values | |
| raan, argp = raan.deg2rad(), argp.deg2rad() | |
| f = torch.stack( | |
| [ | |
| (a / RE).log(), | |
| da.asinh(), | |
| (e + 1e-5).log(), | |
| i / 90, | |
| (100 * di).asinh(), | |
| raan.sin(), | |
| raan.cos(), | |
| argp.sin(), | |
| argp.cos(), | |
| (1e4 * bstar).asinh(), | |
| (1e5 * ndot).asinh(), | |
| ], | |
| -1, | |
| ) | |
| return f.nan_to_num(), n.isnan() | |
| def orbit_targets(x): | |
| """Regime index, median inclination and semi-major axis drift in km/day of each window.""" | |
| n, e, i = x[..., 0], x[..., 1], x[..., 2] | |
| a = semi_major(n) | |
| n_mid, e_mid = n.nanmedian(1).values, e.nanmedian(1).values | |
| apogee = semi_major(n_mid) * (1 + e_mid) - RE | |
| regime = torch.where( | |
| e_mid >= 0.25, 3, torch.where(apogee < 2000, 0, torch.where((n_mid - 1).abs() < 0.1, 2, 1)) | |
| ) | |
| valid = (~n.isnan()).float() | |
| t = torch.arange(x.shape[1], device=x.device).float() | |
| t = t - (valid * t).sum(1, keepdim=True) / valid.sum(1, keepdim=True) | |
| a = (a - a.nanmean(1, keepdim=True)).nan_to_num() | |
| drift = (valid * t * a).sum(1) / (valid * t * t).sum(1).clamp(min=1) | |
| return regime, i.nanmedian(1).values, drift | |
| class Orbit2Vec(nn.Module, PyTorchModelHubMixin): | |
| """Transformer over the daily mean elements of a window, mean-pooled into a unit vector.""" | |
| def __init__( | |
| self, | |
| dim=256, | |
| depth=4, | |
| heads=8, | |
| embed_dim=256, | |
| groups=("other payload",), | |
| owners=("other",), | |
| feature_mean=(0.0,) * FEATURES, | |
| feature_std=(1.0,) * FEATURES, | |
| target_mean=(0.0, 0.0), | |
| target_std=(1.0, 1.0), | |
| ): | |
| super().__init__() | |
| self.groups, self.owners = list(groups), list(owners) | |
| self.feature_mean, self.feature_std = list(feature_mean), list(feature_std) | |
| self.target_mean, self.target_std = list(target_mean), list(target_std) | |
| self.inputs = nn.Linear(FEATURES, dim) | |
| self.positions = nn.Embedding(WINDOW, dim) | |
| layer = nn.TransformerEncoderLayer( | |
| dim, heads, 4 * dim, dropout=0.1, batch_first=True, norm_first=True | |
| ) | |
| self.encoder = nn.TransformerEncoder(layer, depth, enable_nested_tensor=False) | |
| self.norm = nn.LayerNorm(dim) | |
| self.project = nn.Sequential(nn.Linear(dim, dim), nn.GELU(), nn.Linear(dim, embed_dim)) | |
| self.sizes = [len(REGIMES), len(self.groups), len(self.owners), 2] | |
| self.heads = nn.Linear(dim, sum(self.sizes)) | |
| def forward(self, x): | |
| """Embedding and head outputs: regime, group and owner logits, standardized regressions.""" | |
| f, pad = features(x) | |
| mean = torch.tensor(self.feature_mean, device=x.device) | |
| std = torch.tensor(self.feature_std, device=x.device) | |
| f = ((f - mean) / std).masked_fill(pad.unsqueeze(-1), 0) | |
| positions = torch.arange(x.shape[1], device=x.device) | |
| h = self.inputs(f) + self.positions(positions) | |
| h = self.norm(self.encoder(h, src_key_padding_mask=pad)) | |
| keep = (~pad).unsqueeze(-1).to(h.dtype) | |
| pooled = (h * keep).sum(1) / keep.sum(1) | |
| z = functional.normalize(self.project(pooled), dim=-1) | |
| return z, *self.heads(pooled).split(self.sizes, -1) | |
| def run(self, x, batch_size=4096): | |
| """Embeddings and head guesses for windows of shape (n, days, 7), days <= 32.""" | |
| device = next(self.parameters()).device | |
| x = torch.as_tensor(x, dtype=torch.float32) | |
| parts = [self(x[i : i + batch_size].to(device)) for i in range(0, len(x), batch_size)] | |
| z, regime, group, owner, reg = (torch.cat(p).float().cpu() for p in zip(*parts)) | |
| reg = reg * torch.tensor(self.target_std) + torch.tensor(self.target_mean) | |
| return { | |
| "embedding": z, | |
| "regime": regime.argmax(-1), | |
| "group": group.argmax(-1), | |
| "owner": owner.argmax(-1), | |
| "inclination": reg[:, 0], | |
| "drift": reg[:, 1].sinh(), | |
| } | |
| def catalogue(satcat, groups, owners, holdout, seed): | |
| """Objects with a group (constellation or object type) and owner; a random share is unseen.""" | |
| prefix = pl.col("object_name").str.extract(r"^([A-Z]+)") | |
| payloads = satcat.filter(pl.col("object_type") == "PAY").select(prefix.alias("prefix")) | |
| top = payloads.filter(pl.col("prefix") != "OBJECT")["prefix"].value_counts() | |
| top = top.top_k(groups, by="count")["prefix"] | |
| owner_top = satcat["owner"].value_counts().top_k(owners, by="count")["owner"].to_list() | |
| unseen = [i for i in satcat["norad_id"] if random.Random(f"{seed}:{i}").random() < holdout] | |
| return satcat.select( | |
| "norad_id", | |
| name="object_name", | |
| group=pl.when(pl.col("object_type") != "PAY") | |
| .then(pl.col("object_type").replace(NON_PAYLOAD)) | |
| .when(prefix.is_in(top.to_list())) | |
| .then(prefix) | |
| .otherwise(pl.lit("other payload")), | |
| owner=pl.when(pl.col("owner").is_in(owner_top)).then("owner").otherwise(pl.lit("other")), | |
| seen=~pl.col("norad_id").is_in(unseen), | |
| ) | |
| def daily(path, keep): | |
| """The latest TLE per object and UTC day, which also drops repeated epochs.""" | |
| return ( | |
| pl.scan_parquet(path) | |
| .select("norad_id", "epoch", *ELEMENTS) | |
| .filter( | |
| pl.col("norad_id").is_in(keep), | |
| pl.col("mean_motion") > 0.05, | |
| pl.col("eccentricity").is_between(0, 1, closed="left"), | |
| ) | |
| .with_columns(pl.col(ELEMENTS).cast(pl.Float32), day=pl.col("epoch").dt.epoch("d")) | |
| .sort("norad_id", "epoch") | |
| .unique(["norad_id", "day"], keep="last", maintain_order=True) | |
| .drop("epoch") | |
| ) | |
| def gather(days, picks, length): | |
| """Elements on days [start, start + 2 * length) of each pick as halves a and b; NaN: no TLE.""" | |
| grid = ( | |
| picks.select("norad_id", "start") | |
| .with_row_index("row") | |
| .with_columns(offset=pl.int_ranges(0, 2 * length, dtype=pl.Int32)) | |
| .explode("offset") | |
| .with_columns(day=pl.col("start") + pl.col("offset")) | |
| ) | |
| rows = grid.lazy().join(days, on=["norad_id", "day"]).collect() | |
| x = np.full((len(picks), 2 * length, len(ELEMENTS)), np.nan, np.float32) | |
| x[rows["row"].to_numpy(), rows["offset"].to_numpy()] = rows.select(ELEMENTS).to_numpy() | |
| flat = x.reshape(len(picks), 2, -1) | |
| return picks.with_columns(a=pl.Series(flat[:, 0]), b=pl.Series(flat[:, 1])) | |
| def probe(days, labels): | |
| """Ten days before and after each labelled manoeuvre, and around quiet days every 5 days.""" | |
| marks = labels.select("norad_id", day=pl.col("epoch").dt.epoch("d").cast(pl.Int32)).unique() | |
| span = marks.group_by("norad_id").agg(first=pl.col("day").min(), last=pl.col("day").max()) | |
| quiet = ( | |
| span.with_columns(day=pl.int_ranges("first", "last", 5, dtype=pl.Int32)) | |
| .explode("day") | |
| .select("norad_id", "day") | |
| ) | |
| near = quiet.join(marks, on="norad_id").filter( | |
| (pl.col("day") - pl.col("day_right")).abs() <= PROBE_DAYS | |
| ) | |
| picks = pl.concat( | |
| [ | |
| marks.with_columns(maneuver=True), | |
| quiet.join(near, on=["norad_id", "day"], how="anti").with_columns(maneuver=False), | |
| ] | |
| ).sort("norad_id", "day") | |
| picks = gather(days, picks.with_columns(start=pl.col("day") - PROBE_DAYS), PROBE_DAYS) | |
| valid = [(~np.isnan(picks[h].to_numpy()[:, ::7])).sum(1) >= 5 for h in ("a", "b")] | |
| return picks.filter(pl.Series(valid[0] & valid[1])).drop("start") | |
| def data(args): | |
| """Build daily mean-element windows, the object table and the manoeuvre probe.""" | |
| start = time.time() | |
| satcat = pl.read_parquet(hf_hub_download(SATCAT, "data/satcat.parquet", repo_type="dataset")) | |
| labels = pl.read_parquet( | |
| hf_hub_download(MANEUVERS, "maneuver_labels/data.parquet", repo_type="dataset") | |
| ) | |
| objects = catalogue(satcat, args.groups, args.owners, args.holdout, args.seed) | |
| if args.limit: | |
| ids = random.Random(args.seed).sample(sorted(objects["norad_id"]), args.limit) | |
| objects = objects.filter(pl.col("norad_id").is_in(ids + labels["norad_id"].to_list())) | |
| files = sorted( | |
| f | |
| for f in HfApi().list_repo_files(TLE, repo_type="dataset") | |
| if f.startswith("data/tle_") and int(f[9:13]) >= args.first_year | |
| ) | |
| folder = Path(snapshot_download(TLE, repo_type="dataset", allow_patterns=files)) | |
| out = Path(args.output) | |
| (out / "daily").mkdir(parents=True, exist_ok=True) | |
| for f in files: | |
| daily(folder / f, objects["norad_id"].to_list()).sink_parquet(out / "daily" / Path(f).name) | |
| print(json.dumps({"file": f, "minutes": round((time.time() - start) / 60, 1)}), flush=True) | |
| days = pl.scan_parquet(out / "daily" / "*.parquet") | |
| valid = ( | |
| days.group_by("norad_id", window=pl.col("day") // WINDOW) | |
| .len() | |
| .filter(pl.col("len") >= MIN_DAYS) | |
| .drop("len") | |
| .collect() | |
| ) | |
| pairs = ( | |
| valid.join(valid.with_columns(pl.col("window") - 1), on=["norad_id", "window"]) | |
| .join(objects.select("norad_id", "seen"), on="norad_id") | |
| .sort("norad_id", "window") | |
| .with_columns(start=pl.col("window") * WINDOW) | |
| ) | |
| seen = pairs.filter( | |
| pl.col("seen"), pl.col("window") + 1 < TRAIN_END, pl.col("window") % 2 == 0 | |
| ).filter(pl.int_range(pl.len()).shuffle(args.seed).over("norad_id") < args.pairs) | |
| tables = { | |
| ("objects", "train"): objects, | |
| ("windows", "train"): gather(days, seen, WINDOW).drop("seen", "start"), | |
| ("windows", "test"): gather(days, pairs.filter(window=TEST), WINDOW).drop("seen", "start"), | |
| ("maneuvers", "test"): probe(days, labels), | |
| } | |
| for (config, split), part in tables.items(): | |
| part.write_parquet(out / f"{config}-{split}.parquet") | |
| stats = {f"{config}-{split}": len(part) for (config, split), part in tables.items()} | |
| stats["train_objects"] = tables["windows", "train"]["norad_id"].n_unique() | |
| stats["probe_maneuvers"] = int(tables["maneuvers", "test"]["maneuver"].sum()) | |
| print(json.dumps({**stats, "minutes": round((time.time() - start) / 60, 1)}), flush=True) | |
| if args.push: | |
| for (config, split), part in tables.items(): | |
| Dataset.from_polars(part).push_to_hub( | |
| args.repo, config_name=config, split=split, private=True | |
| ) | |
| def table(source, config, split): | |
| """A config and split from the Hub or from a local data folder.""" | |
| if Path(source).is_dir(): | |
| return pl.read_parquet(Path(source) / f"{config}-{split}.parquet") | |
| pattern = f"{config}/{split}-*.parquet" | |
| folder = snapshot_download(source, repo_type="dataset", allow_patterns=pattern) | |
| return pl.read_parquet(Path(folder) / pattern) | |
| def windows(frame, column): | |
| """A window column as a float tensor of shape (rows, days, elements).""" | |
| return torch.tensor(frame[column].to_numpy()).reshape(len(frame), -1, len(ELEMENTS)) | |
| def schedule(step, warmup, total): | |
| """Linear warmup, then cosine decay to zero.""" | |
| if step < warmup: | |
| return (step + 1) / warmup | |
| return 0.5 * (1 + math.cos(math.pi * (step - warmup) / max(1, total - warmup))) | |
| def push_result(repo, name, result, revision=None): | |
| """Upload a result as results/<name>.json in the model repo.""" | |
| HfApi().upload_file( | |
| path_or_fileobj=json.dumps(result, indent=1).encode(), | |
| path_in_repo=f"results/{name}.json", | |
| repo_id=repo, | |
| revision=revision, | |
| commit_message=f"Add {name} results", | |
| ) | |
| def augment(x, min_days, dropout): | |
| """A random crop of min_days..32 days with some days dropped; keeps at least one TLE.""" | |
| rows, length = x.shape[:2] | |
| present = ~x[..., 0].isnan() | |
| size = torch.randint(min_days, length + 1, (rows, 1), device=x.device) | |
| begin = (torch.rand(rows, 1, device=x.device) * (length - size + 1)).floor() | |
| t = torch.arange(length, device=x.device) | |
| keep = (t >= begin) & (t < begin + size) & (torch.rand(rows, length, device=x.device) > dropout) | |
| keep &= present | |
| empty = keep.sum(1) == 0 | |
| keep[empty] = present[empty] | |
| return x.masked_fill(~keep.unsqueeze(-1), math.nan) | |
| def train(args): | |
| """Contrastive training on adjacent windows of the same object, plus orbit heads.""" | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| random.seed(args.seed) | |
| np.random.seed(args.seed) | |
| torch.manual_seed(args.seed) | |
| objects = table(args.data, "objects", "train") | |
| pairs = table(args.data, "windows", "train").join(objects, on="norad_id") | |
| groups, owners = sorted(objects["group"].unique()), sorted(objects["owner"].unique()) | |
| first, second = windows(pairs, "a").to(device), windows(pairs, "b").to(device) | |
| labels = torch.tensor( | |
| np.stack([pairs["group"].replace_strict(groups, list(range(len(groups)))).to_numpy(), | |
| pairs["owner"].replace_strict(owners, list(range(len(owners)))).to_numpy()], 1), | |
| device=device, | |
| ) | |
| sample = first[torch.randperm(len(first), device=device)[:20_000]] | |
| f, pad = features(sample) | |
| f = f[~pad] | |
| regime, inclination, drift = orbit_targets(torch.cat([first, second])) | |
| targets = torch.stack([inclination, drift.asinh()], 1) | |
| target_mean, target_std = targets.mean(0), targets.std(0) + 1e-6 | |
| targets = ((targets - target_mean) / target_std).reshape(2, len(first), 2) | |
| regime = regime.reshape(2, len(first)) | |
| model = Orbit2Vec( | |
| dim=args.dim, | |
| depth=args.depth, | |
| groups=groups, | |
| owners=owners, | |
| feature_mean=f.mean(0).tolist(), | |
| feature_std=(f.std(0) + 1e-6).tolist(), | |
| target_mean=target_mean.tolist(), | |
| target_std=target_std.tolist(), | |
| ).to(device) | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.05) | |
| steps = args.max_steps | |
| lr_schedule = torch.optim.lr_scheduler.LambdaLR( | |
| optimizer, lambda step: schedule(step, min(args.warmup, steps // 10 + 1), steps) | |
| ) | |
| ids, window = pairs["norad_id"].to_numpy(), pairs["window"].to_numpy() | |
| by_window = {w: np.flatnonzero(window == w) for w in np.unique(window)} | |
| batch_size = min(args.batch_size, pairs["norad_id"].n_unique()) | |
| print(json.dumps({"pairs": len(pairs), "objects": pairs["norad_id"].n_unique(), | |
| "parameters": sum(p.numel() for p in model.parameters())}), flush=True) | |
| start, log = time.time(), {} | |
| model.train() | |
| for step in range(steps): | |
| # Most of a batch shares one window, so siblings in the same shell compete. | |
| anchor = by_window[window[random.randrange(len(window))]] | |
| picks = np.concatenate( | |
| [np.random.permutation(anchor), np.random.randint(0, len(window), 4 * batch_size)] | |
| ) | |
| _, firsts = np.unique(ids[picks], return_index=True) | |
| batch = torch.tensor(picks[np.sort(firsts)][:batch_size], device=device) | |
| views = torch.cat([augment(part[batch], args.min_days, args.dropout) | |
| for part in (first, second)]) | |
| with torch.autocast(device, dtype=torch.bfloat16, enabled=device == "cuda"): | |
| z, regime_logits, group_logits, owner_logits, guess = model(views) | |
| z = z.float() | |
| n = len(batch) | |
| similarity = z[:n] @ z[n:].T / args.temperature | |
| gold = torch.arange(n, device=device) | |
| contrastive = (functional.cross_entropy(similarity, gold) | |
| + functional.cross_entropy(similarity.T, gold)) / 2 | |
| heads = ( | |
| functional.cross_entropy(regime_logits.float(), regime[:, batch].reshape(-1)) | |
| + functional.cross_entropy(group_logits.float(), labels[batch, 0].repeat(2)) | |
| + functional.cross_entropy(owner_logits.float(), labels[batch, 1].repeat(2)) | |
| + functional.mse_loss(guess.float(), targets[:, batch].reshape(-1, 2)) | |
| ) | |
| loss = contrastive + args.aux_weight * heads | |
| optimizer.zero_grad(set_to_none=True) | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| optimizer.step() | |
| lr_schedule.step() | |
| if step % 100 == 0 or step == steps - 1: | |
| log = { | |
| "step": step, | |
| "loss": round(loss.item(), 4), | |
| "contrastive": round(contrastive.item(), 4), | |
| "heads": round(heads.item(), 4), | |
| "pair_top1": round((similarity.argmax(1) == gold).float().mean().item(), 3), | |
| "minutes": round((time.time() - start) / 60, 1), | |
| } | |
| print(json.dumps(log), flush=True) | |
| model.eval() | |
| model.save_pretrained(args.output) | |
| (Path(args.output) / "README.md").unlink(missing_ok=True) | |
| if args.push: | |
| api = HfApi() | |
| if args.revision: | |
| api.create_branch(args.repo, branch=args.revision, exist_ok=True) | |
| api.upload_folder( | |
| folder_path=args.output, | |
| repo_id=args.repo, | |
| revision=args.revision, | |
| commit_message="Upload model", | |
| ) | |
| api.upload_file( | |
| path_or_fileobj=__file__, | |
| path_in_repo="orbit2vec.py", | |
| repo_id=args.repo, | |
| revision=args.revision, | |
| ) | |
| push_result( | |
| args.repo, | |
| "train", | |
| { | |
| **log, | |
| "data": args.data, | |
| "pairs": len(pairs), | |
| "objects": pairs["norad_id"].n_unique(), | |
| "steps": steps, | |
| "batch_size": batch_size, | |
| "learning_rate": args.lr, | |
| "temperature": args.temperature, | |
| "aux_weight": args.aux_weight, | |
| "dim": args.dim, | |
| "depth": args.depth, | |
| "parameters": sum(p.numel() for p in model.parameters()), | |
| "runtime_s": round(time.time() - start), | |
| "device": torch.cuda.get_device_name() if device == "cuda" else "cpu", | |
| }, | |
| args.revision, | |
| ) | |
| def nearest(queries, gallery, k=5, chunk=4096): | |
| """Indices of the k most similar gallery rows for each query.""" | |
| k = min(k, len(gallery)) | |
| starts = range(0, len(queries), chunk) | |
| return torch.cat([(queries[i : i + chunk] @ gallery.T).topk(k).indices for i in starts]) | |
| def auc(score, label): | |
| """Chance that a random positive scores above a random negative.""" | |
| ranks = score.argsort().argsort().astype(float) + 1 | |
| positives, negatives = label.sum(), (~label).sum() | |
| return (ranks[label].sum() - positives * (positives + 1) / 2) / (positives * negatives) | |
| def shell_mates(gallery, queries, earlier, later, groups, km=5, degrees=0.5, chunk=1024): | |
| """Top-1 of each later window among earlier ones of the same group, altitude and inclination. | |
| Only queries with at least five such mates count: shell mates look alike, so | |
| average retrieval can hide a model that stops telling them apart. | |
| """ | |
| group = torch.tensor(pl.Series(groups).cast(pl.Categorical).to_physical().to_numpy()) | |
| a = [semi_major(x[..., 0]).nanmedian(1).values for x in (earlier, later)] | |
| i = [x[..., 2].nanmedian(1).values for x in (earlier, later)] | |
| counted, hits = [], [] | |
| for start in range(0, len(later), chunk): | |
| rows = slice(start, start + chunk) | |
| near = ( | |
| (group[rows, None] == group[None]) | |
| & ((a[1][rows, None] - a[0][None]).abs() < km) | |
| & ((i[1][rows, None] - i[0][None]).abs() < degrees) | |
| ) | |
| near[torch.arange(len(near)), torch.arange(start, start + len(near))] = True | |
| top = (queries[rows] @ gallery.T).masked_fill(~near, -2).argmax(1) | |
| counted.append(near.sum(1) >= 5) | |
| hits.append(top == torch.arange(start, start + len(top))) | |
| counted, hits = torch.cat(counted), torch.cat(hits) | |
| return { | |
| "n": int(counted.sum()), | |
| "top1": round(hits[counted].float().mean().item(), 3) if counted.any() else None, | |
| } | |
| def orbit_map(embedding, groups, path): | |
| """UMAP of the gallery's embeddings, coloured by the most common groups.""" | |
| xy = umap.UMAP(metric="cosine", random_state=0).fit_transform(embedding.numpy()) | |
| groups = np.array(groups) | |
| top = [g for g, _ in Counter(groups).most_common(9)] | |
| fig, ax = plt.subplots(figsize=(8, 7)) | |
| rest = ~np.isin(groups, top) | |
| ax.scatter(*xy[rest].T, s=2, linewidths=0, color="0.75", label="other") | |
| for k, group in enumerate(top): | |
| ax.scatter(*xy[groups == group].T, s=2, linewidths=0, color=f"C{k}", label=group) | |
| ax.legend(markerscale=5, fontsize=8, loc="best", frameon=False) | |
| ax.set_axis_off() | |
| fig.savefig(path, dpi=150, bbox_inches="tight") | |
| plt.close(fig) | |
| def evaluate(args): | |
| """Find each object's later window among all earlier ones; check the heads; probe manoeuvres.""" | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| model = Orbit2Vec.from_pretrained(args.model, revision=args.revision).to(device).eval() | |
| objects = table(args.data, "objects", "train") | |
| test = table(args.data, "windows", "test").join(objects, on="norad_id") | |
| if args.limit: | |
| test = test.sample(min(args.limit, len(test)), seed=0) | |
| earlier, later = windows(test, "a"), windows(test, "b") | |
| gallery, out = model.run(earlier), model.run(later) | |
| top = nearest(out["embedding"], gallery["embedding"]) | |
| gold = torch.arange(len(test)) | |
| hits1, hits5 = top[:, 0] == gold, (top == gold[:, None]).any(1) | |
| # Baseline: the mean of the standardized inputs over each window. | |
| mean, std = torch.tensor(model.feature_mean), torch.tensor(model.feature_std) | |
| plain = [] | |
| for x in (earlier, later): | |
| f, pad = features(x) | |
| f = ((f - mean) / std).masked_fill(pad.unsqueeze(-1), math.nan).nanmean(1) | |
| plain.append(functional.normalize(f, dim=-1)) | |
| base1 = nearest(plain[1], plain[0], 1)[:, 0] == gold | |
| regime, inclination, drift = orbit_targets(later) | |
| errors = (out["inclination"] - inclination).abs(), (out["drift"] - drift).abs() | |
| seen = torch.tensor(test["seen"].to_numpy()) | |
| truth = {key: torch.tensor([names.index(v) if v in names else -1 for v in test[key]]) | |
| for key, names in (("group", model.groups), ("owner", model.owners))} | |
| result = {"model": args.model, "gallery": len(test), "chance_top1": round(1 / len(test), 6)} | |
| for name, mask in (("seen", seen), ("unseen", ~seen)): | |
| if not mask.any(): | |
| continue | |
| def rate(hit, mask=mask): | |
| return round(hit[mask].float().mean().item(), 3) | |
| result[name] = { | |
| "n": int(mask.sum()), | |
| "top1": rate(hits1), | |
| "top5": rate(hits5), | |
| "baseline_top1": rate(base1), | |
| "regime_accuracy": rate(out["regime"] == regime), | |
| "group_accuracy": rate(out["group"] == truth["group"]), | |
| "owner_accuracy": rate(out["owner"] == truth["owner"]), | |
| "inclination_mae": round(errors[0][mask].mean().item(), 2), | |
| "drift_mae": round(errors[1][mask].mean().item(), 3), | |
| } | |
| shell = test["group"].to_list() | |
| result["shell_mates"] = shell_mates( | |
| gallery["embedding"], out["embedding"], earlier, later, shell | |
| ) | |
| result["shell_mates"]["baseline_top1"] = shell_mates( | |
| plain[0].nan_to_num(), plain[1].nan_to_num(), earlier, later, shell | |
| )["top1"] | |
| probes = table(args.data, "maneuvers", "test") | |
| before, after = windows(probes, "a"), windows(probes, "b") | |
| distance = 1 - (model.run(before)["embedding"] * model.run(after)["embedding"]).sum(1) | |
| jump = (semi_major(after[..., 0]).nanmedian(1).values | |
| - semi_major(before[..., 0]).nanmedian(1).values).abs() | |
| label = probes["maneuver"].to_numpy() | |
| result["maneuvers"] = { | |
| "objects": probes["norad_id"].n_unique(), | |
| "maneuvers": int(label.sum()), | |
| "quiet": int((~label).sum()), | |
| "auc": round(auc(distance.numpy(), label), 3), | |
| "baseline_auc": round(auc(jump.numpy(), label), 3), | |
| } | |
| print(json.dumps(result, indent=1)) | |
| Path(args.output).mkdir(parents=True, exist_ok=True) | |
| orbit_map(gallery["embedding"], test["group"].to_list(), Path(args.output) / "map.png") | |
| if args.push: | |
| push_result(args.repo, "eval", result, args.revision) | |
| HfApi().upload_file( | |
| path_or_fileobj=Path(args.output) / "map.png", | |
| path_in_repo="results/map.png", | |
| repo_id=args.repo, | |
| revision=args.revision, | |
| ) | |
| def read_tle(path): | |
| """The last 32 days of a TLE file as one window, keeping the latest TLE per day.""" | |
| lines = [line for line in Path(path).read_text(encoding="utf-8").splitlines() | |
| if line[:2] in ("1 ", "2 ")] | |
| rows = {} | |
| for first, second in zip(lines[::2], lines[1::2]): | |
| sat = Satrec.twoline2rv(first, second) | |
| epoch = sat.jdsatepoch + sat.jdsatepochF - 2440587.5 | |
| angles = np.degrees([sat.inclo, sat.nodeo, sat.argpo]) | |
| rows[epoch] = [sat.no_kozai * 720 / math.pi, sat.ecco, *angles, | |
| sat.bstar, sat.ndot * 1440**2 / (2 * math.pi)] | |
| last = int(max(rows)) | |
| x = np.full((1, WINDOW, len(ELEMENTS)), np.nan, np.float32) | |
| for epoch in sorted(rows): | |
| if int(epoch) > last - WINDOW: | |
| x[0, int(epoch) - last + WINDOW - 1] = rows[epoch] | |
| return x | |
| def embed(args): | |
| """Print the nearest catalogued objects and head guesses for one orbit history.""" | |
| model = Orbit2Vec.from_pretrained(args.model).eval() | |
| test = table(args.data, "windows", "test").join( | |
| table(args.data, "objects", "train"), on="norad_id" | |
| ) | |
| if args.tle: | |
| x = read_tle(args.tle) | |
| else: | |
| x = windows(test.filter(norad_id=args.norad), "b") | |
| if x.numel() == 0: | |
| raise SystemExit(f"{args.norad} has no test window; pass --tle") | |
| out = model.run(x) | |
| gallery = model.run(windows(test, "a"))["embedding"] | |
| top = nearest(out["embedding"], gallery, args.top)[0] | |
| similarity = out["embedding"] @ gallery[top].T | |
| print( | |
| json.dumps( | |
| { | |
| "nearest": [ | |
| {"norad_id": test["norad_id"][int(i)], "name": test["name"][int(i)], | |
| "cosine": round(s.item(), 3)} | |
| for i, s in zip(top, similarity[0]) | |
| ], | |
| "regime": REGIMES[out["regime"][0]], | |
| "group": model.groups[out["group"][0]], | |
| "owner": model.owners[out["owner"][0]], | |
| "inclination": round(out["inclination"][0].item(), 2), | |
| "drift_km_per_day": round(out["drift"][0].item(), 3), | |
| "embedding": [round(v, 4) for v in out["embedding"][0].tolist()], | |
| }, | |
| indent=1, | |
| ) | |
| ) | |
| def card(args): | |
| """Render card.jinja into card/README.md with the results stored in the model repo.""" | |
| try: | |
| folder = Path(snapshot_download(args.repo, allow_patterns="results/*.json")) | |
| paths = folder.glob("results/*.json") | |
| except RepositoryNotFoundError: | |
| paths = [] | |
| results = {path.stem: json.loads(path.read_text(encoding="utf-8")) for path in paths} | |
| meta = ModelCardData( | |
| model_name=args.repo.split("/")[1], | |
| datasets=[DATA, TLE, SATCAT, MANEUVERS], | |
| license="mit", | |
| library_name="pytorch", | |
| pipeline_tag="feature-extraction", | |
| tags=["space", "satellites", "tle", "embeddings", "contrastive-learning", "weird2vec"], | |
| ) | |
| rendered = ModelCard.from_template( | |
| meta, | |
| template_path=HERE / "card.jinja", | |
| repo=args.repo, | |
| data=DATA, | |
| tle=TLE, | |
| satcat=SATCAT, | |
| maneuvers=MANEUVERS, | |
| train=results.get("train"), | |
| eval=results.get("eval"), | |
| ) | |
| (HERE / "card").mkdir(exist_ok=True) | |
| rendered.save(HERE / "card" / "README.md") | |
| def main(): | |
| """Parse arguments and run a command.""" | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| commands = parser.add_subparsers(dest="command", required=True) | |
| data_parser = commands.add_parser("data") | |
| data_parser.add_argument("--limit", type=int, default=None, help="number of objects") | |
| data_parser.add_argument("--first-year", type=int, default=1959) | |
| data_parser.add_argument("--pairs", type=int, default=8, help="training pairs per object") | |
| data_parser.add_argument("--holdout", type=float, default=0.1, help="share of unseen objects") | |
| data_parser.add_argument("--groups", type=int, default=24, help="payload name prefixes kept") | |
| data_parser.add_argument("--owners", type=int, default=12, help="owners kept") | |
| data_parser.add_argument("--seed", type=int, default=0) | |
| data_parser.add_argument("--output", default="out/data") | |
| data_parser.add_argument("--repo", default=DATA) | |
| data_parser.add_argument("--push", action="store_true") | |
| data_parser.set_defaults(run=data) | |
| train_parser = commands.add_parser("train") | |
| train_parser.add_argument("--data", default=DATA, help="dataset repo or local data folder") | |
| train_parser.add_argument("--max-steps", type=int, default=20_000) | |
| train_parser.add_argument("--batch-size", type=int, default=512, help="objects per step") | |
| train_parser.add_argument("--lr", type=float, default=3e-4) | |
| train_parser.add_argument("--warmup", type=int, default=1000) | |
| train_parser.add_argument("--temperature", type=float, default=0.05) | |
| train_parser.add_argument("--aux-weight", type=float, default=0.1) | |
| train_parser.add_argument("--min-days", type=int, default=8, help="shortest crop") | |
| train_parser.add_argument("--dropout", type=float, default=0.2, help="share of days dropped") | |
| train_parser.add_argument("--dim", type=int, default=256) | |
| train_parser.add_argument("--depth", type=int, default=4) | |
| train_parser.add_argument("--seed", type=int, default=0) | |
| train_parser.add_argument("--output", default="out/model") | |
| train_parser.add_argument("--repo", default=REPO) | |
| train_parser.add_argument("--revision", default=None, help="branch to push to") | |
| train_parser.add_argument("--push", action="store_true") | |
| train_parser.set_defaults(run=train) | |
| eval_parser = commands.add_parser("eval") | |
| eval_parser.add_argument("--model", default=REPO, help="model repo or local folder") | |
| eval_parser.add_argument("--data", default=DATA) | |
| eval_parser.add_argument("--limit", type=int, default=None, help="number of test objects") | |
| eval_parser.add_argument("--repo", default=REPO, help="where --push stores the result") | |
| eval_parser.add_argument("--revision", default=None, help="model branch") | |
| eval_parser.add_argument("--output", default="out/eval", help="where the orbit map goes") | |
| eval_parser.add_argument("--push", action="store_true") | |
| eval_parser.set_defaults(run=evaluate) | |
| embed_parser = commands.add_parser("embed") | |
| source = embed_parser.add_mutually_exclusive_group(required=True) | |
| source.add_argument("--norad", type=int, help="catalogue number of an object in the test set") | |
| source.add_argument("--tle", help="text file of TLEs for one object, oldest first") | |
| embed_parser.add_argument("--model", default=REPO) | |
| embed_parser.add_argument("--data", default=DATA) | |
| embed_parser.add_argument("--top", type=int, default=5) | |
| embed_parser.set_defaults(run=embed) | |
| card_parser = commands.add_parser("card") | |
| card_parser.add_argument("--repo", default=REPO) | |
| card_parser.set_defaults(run=card) | |
| args = parser.parse_args() | |
| torch.set_num_threads(cpus()) | |
| if not PROGRESS: | |
| datasets.disable_progress_bars() | |
| huggingface_hub.utils.disable_progress_bars() | |
| args.run(args) | |
| if __name__ == "__main__": | |
| main() | |