orbitsight / scripts /make_cpu_student_tracking_animations.py
HishaamA's picture
Add DAVIS clip to animation renderer
a06face verified
Raw History Blame Contribute Delete
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"
@dataclass(frozen=True)
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())