# /// 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/.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()