orbit2vec / orbit2vec.py
jgalego's picture
promote aux0.3
90e7c9b verified
Raw History Blame Contribute Delete
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)
@torch.no_grad()
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()