arissassina's picture
download
raw
2.42 kB
import gc
import json
import os
import cv2
import numpy as np
import torch
from cotracker.predictor import CoTrackerPredictor
ROOT = "/root/bdmc_pipeline"
H, W = 1920, 1080
WIN_LEN = 24
STRIDE = 12
GRID_STEP = 44
Z_MIN, Z_MAX = 1.2, 9.0
CKPT = f"{ROOT}/models/cotracker3/scaled_offline.pth"
out_dir = f"{ROOT}/outputs/tracks"
os.makedirs(out_dir, exist_ok=True)
print("Loading CoTracker3...")
model = CoTrackerPredictor(checkpoint=CKPT).to("cuda").eval()
def load_window(f0, n):
frames = []
for i in range(f0, min(301, f0 + n)):
im = cv2.imread(f"{ROOT}/frames/f_{i:04d}.jpg")
frames.append(cv2.cvtColor(im, cv2.COLOR_BGR2RGB))
vid = torch.from_numpy(np.stack(frames)).permute(0, 3, 1, 2)[None].float() / 255.0
return vid.to("cuda")
summary = []
w_idx = 0
for f0 in range(1, 301 - 15, STRIDE):
n = min(WIN_LEN, 301 - f0)
if n < 16:
break
dep = np.load(f"{ROOT}/depth/f_{f0:04d}.npy").astype(np.float32)
m = cv2.imread(f"{ROOT}/masks/road_surface/f_{f0:04d}.png", 0) > 127
valid = m & (dep > Z_MIN) & (dep < Z_MAX)
ys, xs = np.nonzero(valid)
sel = (ys % GRID_STEP == 0) & (xs % GRID_STEP == 0)
ys, xs = ys[sel], xs[sel]
if len(ys) > 420:
idx = np.random.default_rng(f0).choice(len(ys), 420, replace=False)
ys, xs = ys[idx], xs[idx]
queries = torch.tensor(np.stack([np.zeros_like(xs), xs, ys], 1), dtype=torch.float32)[None].to("cuda")
video = load_window(f0, n)
with torch.no_grad():
out = model(video, queries=queries, backward_tracking=False)
tracks, vis = out[0], out[1]
tracks = tracks[0].cpu().numpy()
vis = vis[0].cpu().numpy().astype(bool)
np.savez_compressed(f"{out_dir}/win_{f0:04d}.npz", tracks=tracks, vis=vis,
start=f0, n_frames=n, xs=xs, ys=ys)
end_vis = float(vis[-1].mean())
summary.append({"window_start": f0, "frames": n, "n_points": int(len(xs)),
"end_visibility_frac": round(end_vis, 3)})
print(f"win {f0}+{n}: {len(xs)} pts, end-vis {end_vis:.2f}", flush=True)
del video
gc.collect()
torch.cuda.empty_cache()
w_idx += 1
with open(f"{ROOT}/outputs/phase5_summary.json", "w") as fj:
json.dump({"windows": summary,
"mean_end_visibility": round(float(np.mean([s["end_visibility_frac"] for s in summary])), 3)}, fj, indent=1)
print("Phase 5a tracking complete", flush=True)

Xet Storage Details

Size:
2.42 kB
·
Xet hash:
7c64c39a1f18173551f54e24d0cdd1c934ee0a58c2cd1f7e428920770c0b9a37

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.