Download scripts/make_cpu_student_tracking_animations.py from simerai/orbitsight: direct link, hf CLI and curl.
- Browser
- Download file 13.9 kB
-
https://huggingface.co/simerai/orbitsight/resolve/main/scripts/make_cpu_student_tracking_animations.py
- Command line
-
hf download hf://simerai/orbitsight/scripts/make_cpu_student_tracking_animations.py
-
curl -L -o make_cpu_student_tracking_animations.py https://huggingface.co/simerai/orbitsight/resolve/main/scripts/make_cpu_student_tracking_animations.py
13.9 kB
| #!/usr/bin/env python3 | |
| """Create real-data GIFs of the promoted CPU student's detector and tracker.""" | |
| from __future__ import annotations | |
| import hashlib | |
| import json | |
| import pickle | |
| import sys | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| import numpy as np | |
| from PIL import Image, ImageDraw, ImageFont | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT / "src")) | |
| from orbitsight.eval.metrics import iou_xywh # noqa: E402 | |
| from orbitsight.ingestion import discover_sequences, iter_windows, load_events, load_gt_boxes # noqa: E402 | |
| from orbitsight.postprocessing import finalize_detections, rank_detections # noqa: E402 | |
| from orbitsight.tracking import get_tracking_policy # noqa: E402 | |
| from orbitsight.tracking.tracker import tracks_to_detections # noqa: E402 | |
| from orbitsight.viz.visualize import _event_image # noqa: E402 | |
| CACHE = ROOT / "artifacts/cpu_student/post_training_latency_neutral/current_student_top1_cache.pkl" | |
| OUT = ROOT / "competition_ready_improved_cpu_student/visualizations/animations" | |
| BG = "#07111f" | |
| PANEL = "#0e2035" | |
| TEXT = "#e9f2f9" | |
| MUTED = "#91a7ba" | |
| DETECTOR = "#ff5a5f" | |
| TRACKER = "#39e6ff" | |
| GT = "#ffd43b" | |
| GOOD = "#54e38e" | |
| class Clip: | |
| slug: str | |
| sequence: str | |
| start: int | |
| end: int | |
| title: str | |
| subtitle: str | |
| CLIPS = ( | |
| Clip( | |
| "davis_saocom_tracker_smoothing", | |
| "DAVIS_SAOCOM1B_46265_2024-12-04-18-21-37", | |
| 3388, | |
| 3416, | |
| "DAVIS SAOCOM-1B: detector + tracker", | |
| "Native 346×260 events; cyan smoothing stabilizes the raw red localization", | |
| ), | |
| Clip( | |
| "evk4_bright_track", | |
| "2025_12_23_20_53_46_EVK4_mag7.3", | |
| 82, | |
| 113, | |
| "EVK4: detector + recurrent tracker", | |
| "Large-sensor event field; cyan Kalman state follows the raw red proposal", | |
| ), | |
| Clip( | |
| "dvx_stars_tracker_smoothing", | |
| "DVX_Filtered_Stars3_2025-01-20-20-22-53", | |
| 4136, | |
| 4165, | |
| "DVX Stars3: tracker smoothing in clutter", | |
| "A real interval where tracking materially improves several localization overlaps", | |
| ), | |
| Clip( | |
| "dvx_thuraya_low_snr", | |
| "DVX_Filtered_Thuraya3_32404_2025-01-20-20-02-43", | |
| 2445, | |
| 2477, | |
| "DVX Thuraya3: low-SNR target", | |
| "The hardest regime: tiny boxes, weak event evidence, and changing confidence", | |
| ), | |
| ) | |
| def font(size: int, bold: bool = False): | |
| name = "DejaVuSans-Bold.ttf" if bold else "DejaVuSans.ttf" | |
| path = Path("/usr/share/fonts/truetype/dejavu") / name | |
| return ImageFont.truetype(str(path), size) if path.exists() else ImageFont.load_default() | |
| def dashed_rectangle(draw: ImageDraw.ImageDraw, box, fill, width=3, dash=8): | |
| x1, y1, x2, y2 = [int(round(v)) for v in box] | |
| for x in range(x1, x2, dash * 2): | |
| draw.line((x, y1, min(x + dash, x2), y1), fill=fill, width=width) | |
| draw.line((x, y2, min(x + dash, x2), y2), fill=fill, width=width) | |
| for y in range(y1, y2, dash * 2): | |
| draw.line((x1, y, x1, min(y + dash, y2)), fill=fill, width=width) | |
| draw.line((x2, y, x2, min(y + dash, y2)), fill=fill, width=width) | |
| def xyxy(box): | |
| cx, cy, w, h = box[:4] | |
| return cx - w / 2, cy - h / 2, cx + w / 2, cy + h / 2 | |
| def replay_diagnostics(row: dict, end: int) -> dict[int, dict]: | |
| policy = get_tracking_policy("cpu_student") | |
| tracker = policy.build_tracker() | |
| records = {} | |
| for index, detections in enumerate(row["detections"][: end + 1]): | |
| fed = rank_detections( | |
| [tuple(float(v) for v in det) for det in detections | |
| if float(det[4]) >= policy.detector_threshold], | |
| policy.pre_max_det, | |
| ) | |
| tracks = tracker.update(fed) | |
| emitted = tracks_to_detections( | |
| tracks, | |
| emit_coasting=policy.emit_coasting, | |
| max_coast_age=policy.max_coast_age, | |
| coast_decay=policy.coast_decay, | |
| ) | |
| emitted = finalize_detections(emitted, row["sensor_profile"], policy.post_max_det) | |
| emitted = [det for det in emitted if det[4] >= row["emit_gate"]] | |
| chosen_id = None | |
| chosen_age = None | |
| if emitted: | |
| target = emitted[0] | |
| nearest = min( | |
| tracks, | |
| key=lambda track: (track.box[0] - target[0]) ** 2 + (track.box[1] - target[1]) ** 2, | |
| ) | |
| chosen_id, chosen_age = nearest.id, nearest.age | |
| records[index] = { | |
| "detector": fed[:1], | |
| "tracker": emitted[:1], | |
| "track_id": chosen_id, | |
| "track_age": chosen_age, | |
| } | |
| return records | |
| def selected_windows(ref, start: int, end: int): | |
| selected = {} | |
| events = load_events(ref, mmap=True) | |
| gt = load_gt_boxes(ref) | |
| for window in iter_windows(events, ref.name, gt): | |
| if start <= window.index <= end: | |
| selected[window.index] = window | |
| if window.index >= end: | |
| break | |
| expected = set(range(start, end + 1)) | |
| if set(selected) != expected: | |
| raise RuntimeError(f"missing animation windows for {ref.name}: {sorted(expected-set(selected))}") | |
| return selected | |
| def crop_bounds(windows, records, sensor, start, end): | |
| points = [] | |
| for index in range(start, end + 1): | |
| window = windows[index] | |
| for gt in window.gt_boxes: | |
| points.append((gt.cx, gt.cy)) | |
| for key in ("detector", "tracker"): | |
| for box in records[index][key]: | |
| points.append((box[0], box[1])) | |
| if not points: | |
| return 0, 0, sensor.width, sensor.height | |
| xs, ys = zip(*points) | |
| span = max(max(xs) - min(xs), max(ys) - min(ys), 95 if sensor.name == "DVX" else 155) | |
| cx, cy = (min(xs) + max(xs)) / 2, (min(ys) + max(ys)) / 2 | |
| half = span * 0.72 | |
| x1, x2 = max(0, cx - half), min(sensor.width, cx + half) | |
| y1, y2 = max(0, cy - half), min(sensor.height, cy + half) | |
| return x1, y1, x2, y2 | |
| def fit_rect(src_w, src_h, dst): | |
| x, y, w, h = dst | |
| scale = min(w / src_w, h / src_h) | |
| rw, rh = src_w * scale, src_h * scale | |
| return x + (w - rw) / 2, y + (h - rh) / 2, rw, rh | |
| def map_box(box, source_bounds, target_rect): | |
| sx1, sy1, sx2, sy2 = source_bounds | |
| tx, ty, tw, th = target_rect | |
| x1, y1, x2, y2 = xyxy(box) | |
| return ( | |
| tx + (x1 - sx1) / (sx2 - sx1) * tw, | |
| ty + (y1 - sy1) / (sy2 - sy1) * th, | |
| tx + (x2 - sx1) / (sx2 - sx1) * tw, | |
| ty + (y2 - sy1) / (sy2 - sy1) * th, | |
| ) | |
| def render_clip(clip: Clip, ref, row: dict) -> tuple[Path, Path, dict]: | |
| sensor = ref.sensor | |
| windows = selected_windows(ref, clip.start, clip.end) | |
| records = replay_diagnostics(row, clip.end) | |
| crop = crop_bounds(windows, records, sensor, clip.start, clip.end) | |
| canvas_size = (960, 540) | |
| full_slot = (20, 72, 565, 440) | |
| zoom_slot = (610, 72, 330, 330) | |
| full_rect = fit_rect(sensor.width, sensor.height, full_slot) | |
| zoom_rect = fit_rect(crop[2] - crop[0], crop[3] - crop[1], zoom_slot) | |
| detector_trail, tracker_trail = [], [] | |
| frames = [] | |
| stats = {"frames": clip.end - clip.start + 1, "tracker_iou_better": 0, "gt_frames": 0} | |
| for index in range(clip.start, clip.end + 1): | |
| window = windows[index] | |
| record = records[index] | |
| gt = [(g.cx, g.cy, g.w, g.h) for g in window.gt_boxes] | |
| detector = record["detector"] | |
| tracked = record["tracker"] | |
| if detector: | |
| detector_trail.append((detector[0][0], detector[0][1])) | |
| if tracked: | |
| tracker_trail.append((tracked[0][0], tracked[0][1])) | |
| base = Image.fromarray(_event_image(np.asarray(window.events), sensor)).convert("RGB") | |
| canvas = Image.new("RGB", canvas_size, BG) | |
| draw = ImageDraw.Draw(canvas) | |
| draw.text((20, 13), clip.title, fill=TEXT, font=font(22, True)) | |
| draw.text((20, 42), clip.subtitle, fill=MUTED, font=font(12)) | |
| full = base.resize((int(full_rect[2]), int(full_rect[3])), Image.Resampling.BILINEAR) | |
| canvas.paste(full, (int(full_rect[0]), int(full_rect[1]))) | |
| crop_img = base.crop(tuple(int(round(v)) for v in crop)) | |
| crop_img = crop_img.resize((int(zoom_rect[2]), int(zoom_rect[3])), Image.Resampling.NEAREST) | |
| canvas.paste(crop_img, (int(zoom_rect[0]), int(zoom_rect[1]))) | |
| draw = ImageDraw.Draw(canvas) | |
| draw.rectangle((full_rect[0], full_rect[1], full_rect[0]+full_rect[2], full_rect[1]+full_rect[3]), outline="#29445e", width=2) | |
| draw.rectangle((zoom_rect[0], zoom_rect[1], zoom_rect[0]+zoom_rect[2], zoom_rect[1]+zoom_rect[3]), outline="#29445e", width=2) | |
| draw.text((zoom_rect[0]+8, zoom_rect[1]+7), "TARGET ZOOM", fill=TEXT, font=font(12, True), stroke_width=2, stroke_fill=BG) | |
| all_bounds = (0, 0, sensor.width, sensor.height) | |
| for trail, color in ((detector_trail[-18:], DETECTOR), (tracker_trail[-18:], TRACKER)): | |
| for bounds, rect in ((all_bounds, full_rect), (crop, zoom_rect)): | |
| mapped = [( | |
| rect[0] + (x-bounds[0])/(bounds[2]-bounds[0])*rect[2], | |
| rect[1] + (y-bounds[1])/(bounds[3]-bounds[1])*rect[3], | |
| ) for x, y in trail if bounds[0] <= x <= bounds[2] and bounds[1] <= y <= bounds[3]] | |
| if len(mapped) > 1: | |
| draw.line(mapped, fill=color, width=2) | |
| for bounds, rect, width in ((all_bounds, full_rect, 2), (crop, zoom_rect, 4)): | |
| for box in gt: | |
| dashed_rectangle(draw, map_box(box, bounds, rect), GT, width=width, dash=8) | |
| for box in detector: | |
| draw.rectangle(map_box(box, bounds, rect), outline=DETECTOR, width=width) | |
| for box in tracked: | |
| draw.rectangle(map_box(box, bounds, rect), outline=TRACKER, width=width) | |
| det_iou = tracker_iou = None | |
| if gt: | |
| stats["gt_frames"] += 1 | |
| if detector: | |
| det_iou = iou_xywh(detector[0][:4], gt[0]) | |
| if tracked: | |
| tracker_iou = iou_xywh(tracked[0][:4], gt[0]) | |
| if det_iou is not None and tracker_iou is not None and tracker_iou > det_iou: | |
| stats["tracker_iou_better"] += 1 | |
| info_y = 418 | |
| draw.rounded_rectangle((610, info_y, 940, 520), radius=9, fill=PANEL, outline="#29445e") | |
| elapsed = (window.t_start_us - windows[clip.start].t_start_us) / 1e6 | |
| draw.text((625, info_y+10), f"window {index} +{elapsed:0.2f}s events {len(window.events):,}", fill=TEXT, font=font(12, True)) | |
| det_text = "none" if not detector else f"{detector[0][4]:.3f} IoU {det_iou:.2f}" if det_iou is not None else f"{detector[0][4]:.3f}" | |
| trk_text = "none" if not tracked else f"#{record['track_id']} age {record['track_age']} IoU {tracker_iou:.2f}" if tracker_iou is not None else f"#{record['track_id']} age {record['track_age']}" | |
| draw.text((625, info_y+36), f"DETECTOR {det_text}", fill=DETECTOR, font=font(12, True)) | |
| draw.text((625, info_y+58), f"TRACKER {trk_text}", fill=TRACKER, font=font(12, True)) | |
| draw.text((625, info_y+80), "GT dashed yellow", fill=GT, font=font(11)) | |
| progress = (index - clip.start + 1) / (clip.end - clip.start + 1) | |
| draw.rectangle((20, 526, 920, 532), fill="#1b334a") | |
| draw.rectangle((20, 526, 20 + int(900 * progress), 532), fill=GOOD) | |
| frames.append(canvas.quantize(colors=128, method=Image.Quantize.FASTOCTREE)) | |
| OUT.mkdir(parents=True, exist_ok=True) | |
| gif = OUT / f"{clip.slug}.gif" | |
| poster = OUT / f"{clip.slug}_poster.png" | |
| frames[0].convert("RGB").save(poster, optimize=True) | |
| frames[0].save(gif, save_all=True, append_images=frames[1:], duration=125, | |
| loop=0, optimize=True, disposal=2) | |
| stats.update({"sequence": clip.sequence, "start_window": clip.start, | |
| "end_window": clip.end, "gif": gif.name, "poster": poster.name}) | |
| return gif, poster, stats | |
| def refresh_checksums(bundle: Path) -> None: | |
| lines = [] | |
| for path in sorted(p for p in bundle.rglob("*") if p.is_file() and p.name != "CHECKSUMS.sha256"): | |
| lines.append(f"{hashlib.sha256(path.read_bytes()).hexdigest()} {path.relative_to(bundle)}") | |
| (bundle / "CHECKSUMS.sha256").write_text("\n".join(lines) + "\n") | |
| def main() -> int: | |
| with CACHE.open("rb") as stream: | |
| cache = pickle.load(stream) | |
| refs = {ref.name: ref for ref in discover_sequences(str(ROOT / "data"))} | |
| outputs = [] | |
| for clip in CLIPS: | |
| gif, _poster, stats = render_clip(clip, refs[clip.sequence], cache[clip.sequence]) | |
| stats["sha256"] = hashlib.sha256(gif.read_bytes()).hexdigest() | |
| outputs.append(stats) | |
| print(f"rendered {gif}") | |
| manifest = { | |
| "source_cache": str(CACHE.relative_to(ROOT)), | |
| "policy": "cpu_student", | |
| "scope": "Exact promoted cached top-1 detector proposals replayed through the final Kalman/coast policy; GT is visualization-only.", | |
| "legend": {"detector": DETECTOR, "tracker": TRACKER, "ground_truth": GT}, | |
| "outputs": outputs, | |
| } | |
| (OUT / "animation_manifest.json").write_text(json.dumps(manifest, indent=2) + "\n") | |
| (OUT / "README.md").write_text( | |
| "# OrbitSight detector + tracker animations\n\n" | |
| "These GIFs use real ChallengeON recordings and the exact promoted CPU-student top-one cache. " | |
| "Red is the raw detector proposal, cyan is the Kalman-smoothed emitted box, yellow dashed is " | |
| "human ground truth, and the colored trails show recent centers. Ground truth is used only for " | |
| "annotation and IoU display.\n\n" + | |
| "\n".join(f"- `{row['gif']}` — `{row['sequence']}`, windows {row['start_window']}–{row['end_window']}" for row in outputs) + "\n" | |
| ) | |
| refresh_checksums(OUT.parents[1]) | |
| print(json.dumps(manifest, indent=2)) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |