Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions tests/test_duplicates.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,20 @@ def test_cross_split_leakage_is_flagged(self) -> None:
report = validate_project(Project.open(root))
self.assertIn("split.near_duplicate_leakage", [issue.code for issue in report.errors])

def test_identical_hashes_expand_to_all_distance_zero_pairs(self) -> None:
# Re-encoded copies share the exact hash; the collapsed LSH path must
# still report every pair, at distance 0, and leave distinct images out.
phashes = {"a": "00000000000000ff", "b": "00000000000000ff", "c": "00000000000000ff", "d": "ffffffffffffff00"}
pairs = near_duplicate_pairs(phashes, threshold=5)
self.assertEqual({(p.asset_a, p.asset_b, p.distance) for p in pairs}, {("a", "b", 0), ("a", "c", 0), ("b", "c", 0)})

def test_close_but_not_identical_hashes_still_pair_across_groups(self) -> None:
# 2 bits apart (<= threshold): the pair must survive the identical-hash
# collapse and carry the true distance.
phashes = {"a": "00000000000000ff", "b": "00000000000000fc"}
pairs = near_duplicate_pairs(phashes, threshold=5)
self.assertEqual([(pairs[0].asset_a, pairs[0].asset_b, pairs[0].distance)], [("a", "b", 2)])

def test_phash_is_persisted_on_import(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
Expand Down
78 changes: 78 additions & 0 deletions tests/test_remote_asset_checks.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
"""Regression tests: fsck/validate must handle cloud-backed (remote) assets.

An asset synced against a cloud target records its object URI as ``path``
(e.g. ``s3://bucket/objects/sha256/...``); such assets have no local file to
open. ``run_fsck`` used to crash on the first remote asset and
``validate_project`` flooded the report with false ``image.unreadable`` errors.
"""

from __future__ import annotations

import tempfile
import unittest
from pathlib import Path

from PIL import Image

from visionpack.core.models import Asset
from visionpack.core.project import Project
from visionpack.formats.yolo import YoloImporter
from visionpack.fsck import run_fsck
from visionpack.validation import validate_project


def _seed_with_remote_asset(root: Path) -> Project:
raw = root / "raw"
(raw / "images").mkdir(parents=True)
(raw / "labels").mkdir(parents=True)
(raw / "classes.txt").write_text("cat\n", encoding="utf-8")
Project.init(root, name="remote-checks")
Image.new("RGB", (30, 30), (10, 20, 30)).save(raw / "images" / "local.png", format="PNG")
(raw / "labels" / "local.txt").write_text("0 0.5 0.5 0.4 0.4\n", encoding="utf-8")
YoloImporter(Project.open(root), raw).run()

project = Project.open(root)
digest = "b" * 64
project.index.upsert_asset(
Asset(
id=f"asset_{digest[:16]}",
sha256=digest,
media_type="image",
path=f"s3://bucket/objects/sha256/{digest[:2]}/{digest[2:4]}/{digest}",
original_path="s3://bucket/imgs/remote.jpg",
width=64,
height=64,
channels=3,
format="jpeg",
size_bytes=1234,
phash="00000000000000ff",
)
)
project.index.save()
return Project.open(root)


class RemoteAssetChecksTest(unittest.TestCase):
def test_fsck_skips_remote_assets_instead_of_crashing(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
project = _seed_with_remote_asset(Path(tmp))
report = run_fsck(project)
self.assertTrue(report.ok)
self.assertEqual(report.checked_assets, 2) # remote asset still counted

def test_fsck_deep_skips_remote_assets(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
project = _seed_with_remote_asset(Path(tmp))
report = run_fsck(project, deep=True)
self.assertTrue(report.ok)

def test_validate_does_not_flag_remote_assets_as_unreadable(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
project = _seed_with_remote_asset(Path(tmp))
report = validate_project(project)
unreadable = [issue for issue in report.issues if issue.code == "image.unreadable"]
self.assertEqual(unreadable, [])


if __name__ == "__main__":
unittest.main()
76 changes: 55 additions & 21 deletions visionpack/duplicates.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from __future__ import annotations

from collections import defaultdict
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass

from visionpack.core.errors import VisionPackError
Expand Down Expand Up @@ -46,18 +47,30 @@ def phash_map(project: Project, assets: list[Asset] | None = None, *, persist: b
"""
assets = assets if assets is not None else project.index.assets()
result: dict[str, str] = {}
backfilled: list[Asset] = []
missing: list[Asset] = []
for asset in assets:
if asset.phash:
result[asset.id] = asset.phash
continue
try:
value = dhash_path(asset.resolved_path(project.root))
except VisionPackError:
continue
result[asset.id] = value
asset.phash = value
backfilled.append(asset)
else:
missing.append(asset)

backfilled: list[Asset] = []
if missing:
# Decoding image headers is I/O-bound and per-asset, so the backfill
# fans out across threads (unreadable/remote assets come back as None).
def compute(asset: Asset) -> str | None:
try:
return dhash_path(asset.resolved_path(project.root))
except VisionPackError:
return None

with ThreadPoolExecutor() as pool:
for asset, value in zip(missing, pool.map(compute, missing), strict=True):
if value is None:
continue
result[asset.id] = value
asset.phash = value
backfilled.append(asset)
if persist and backfilled:
for asset in backfilled:
project.index.upsert_asset(asset)
Expand All @@ -66,30 +79,51 @@ def phash_map(project: Project, assets: list[Asset] | None = None, *, persist: b


def near_duplicate_pairs(phashes: dict[str, str], threshold: int = DEFAULT_THRESHOLD) -> list[NearDuplicatePair]:
"""Find every pair of assets within ``threshold`` bits, via LSH bucketing."""
bands = max(1, threshold + 1)
buckets: dict[tuple[int, int], list[tuple[str, int]]] = defaultdict(list)
"""Find every pair of assets within ``threshold`` bits, via LSH bucketing.

Assets sharing the *exact* hash are collapsed to one value first: batches of
re-encoded near-uniform images (flat backgrounds, calibration plates) all
land on the same hash, and running them through the bucket loop individually
is the pathological O(n^2) case. The LSH comparison count is therefore
bounded by the number of *distinct* hashes, and identical-hash groups expand
to distance-0 pairs directly.
"""
by_value: dict[int, list[str]] = defaultdict(list)
for asset_id, phash in phashes.items():
value = int(phash, 16)
for key in band_keys(value, bands):
buckets[key].append((asset_id, value))
by_value[int(phash, 16)].append(asset_id)

seen: set[tuple[str, str]] = set()
pairs: list[NearDuplicatePair] = []
for ids in by_value.values():
if len(ids) > 1:
ids.sort()
pairs.extend(NearDuplicatePair(a, b, 0) for i, a in enumerate(ids) for b in ids[i + 1 :])

bands = max(1, threshold + 1)
buckets: dict[tuple[int, int], list[int]] = defaultdict(list)
for value in by_value:
for key in band_keys(value, bands):
buckets[key].append(value)

# A pair over the threshold can never come back under it, so it is marked
# seen too — no bucket sharing another band re-XORs the same pair.
seen: set[tuple[int, int]] = set()
for bucket in buckets.values():
if len(bucket) < 2:
continue
for i in range(len(bucket)):
a_id, a_val = bucket[i]
a_val = bucket[i]
for j in range(i + 1, len(bucket)):
b_id, b_val = bucket[j]
key = (a_id, b_id) if a_id < b_id else (b_id, a_id)
b_val = bucket[j]
key = (a_val, b_val) if a_val < b_val else (b_val, a_val)
if key in seen:
continue
seen.add(key)
distance = (a_val ^ b_val).bit_count()
if distance <= threshold:
seen.add(key)
pairs.append(NearDuplicatePair(key[0], key[1], distance))
for a_id in by_value[a_val]:
for b_id in by_value[b_val]:
first, second = (a_id, b_id) if a_id < b_id else (b_id, a_id)
pairs.append(NearDuplicatePair(first, second, distance))
pairs.sort(key=lambda pair: (pair.distance, pair.asset_a, pair.asset_b))
return pairs

Expand Down
50 changes: 33 additions & 17 deletions visionpack/eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,9 +101,13 @@ def _evaluate_detection(project: Project, predictions: PredictionSet, asset_ids:
for class_id in class_ids:
class_preds = sorted(preds.get(class_id, []), key=lambda item: -item[0])
num_gt = gt_counts.get(class_id, 0)
aps = {threshold: _average_precision(_match(class_preds, truths[class_id], threshold), num_gt) for threshold in IOU_THRESHOLDS}
confident = [item for item in class_preds if item[0] >= conf_threshold]
tp = sum(_match(confident, truths[class_id], 0.5))
# Each prediction's IoUs against its image's ground truth are computed
# once and replayed for every threshold, instead of once per threshold.
pred_ious = _pred_ious(class_preds, truths[class_id])
aps = {threshold: _average_precision(_match(class_preds, pred_ious, threshold), num_gt) for threshold in IOU_THRESHOLDS}
confident_indices = [index for index, item in enumerate(class_preds) if item[0] >= conf_threshold]
confident = [class_preds[index] for index in confident_indices]
tp = sum(_match(confident, [pred_ious[index] for index in confident_indices], 0.5))
per_class[class_id] = {
"gt": num_gt,
"predictions": len(class_preds),
Expand All @@ -129,25 +133,37 @@ def _evaluate_detection(project: Project, predictions: PredictionSet, asset_ids:
}


def _match(class_preds: list[tuple[float, str, BBox]], gt_by_asset: dict[str, list[BBox]], iou_threshold: float) -> list[bool]:
def _pred_ious(class_preds: list[tuple[float, str, BBox]], gt_by_asset: dict[str, list[BBox]]) -> list[list[tuple[float, int]]]:
"""For each prediction, its ``(iou, gt_index)`` candidates, best IoU first.

The sort is stable, so ties keep ground-truth enumeration order — matching
then behaves exactly like a per-threshold argmax with ``>`` comparison.
"""
ious: list[list[tuple[float, int]]] = []
for _, asset_id, box in class_preds:
candidates = [(bbox_iou(box, gt_box), index) for index, gt_box in enumerate(gt_by_asset.get(asset_id, []))]
candidates = [item for item in candidates if item[0] > 0.0]
candidates.sort(key=lambda item: -item[0])
ious.append(candidates)
return ious


def _match(class_preds: list[tuple[float, str, BBox]], pred_ious: list[list[tuple[float, int]]], iou_threshold: float) -> list[bool]:
"""Greedy COCO-style matching: each prediction (confidence-ordered) claims the
best still-unmatched ground-truth box in its image. Returns TP flags."""
matched: dict[str, set[int]] = defaultdict(set)
flags: list[bool] = []
for _, asset_id, box in class_preds:
candidates = gt_by_asset.get(asset_id, [])
best_iou, best_index = 0.0, -1
for index, gt_box in enumerate(candidates):
if index in matched[asset_id]:
for (_, asset_id, _), candidates in zip(class_preds, pred_ious, strict=True):
hit = False
for iou, gt_index in candidates:
if iou < iou_threshold:
break # sorted best-first: nothing below qualifies either
if gt_index in matched[asset_id]:
continue
iou = bbox_iou(box, gt_box)
if iou > best_iou:
best_iou, best_index = iou, index
if best_index >= 0 and best_iou >= iou_threshold:
matched[asset_id].add(best_index)
flags.append(True)
else:
flags.append(False)
matched[asset_id].add(gt_index)
hit = True
break
flags.append(hit)
return flags


Expand Down
5 changes: 5 additions & 0 deletions visionpack/fsck.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,11 @@ def run_fsck(project: Project, deep: bool = False, check_orphans: bool = True) -
checked_assets += 1
asset_ids.add(asset.id)
referenced_sha.add(asset.sha256)
if asset.is_remote:
# Bytes live in a remote object store; there is no local file to
# check (and resolved_path would raise). Referential checks below
# still cover the asset.
continue
path = asset.resolved_path(project.root)
if not path.exists():
issues.append(FsckIssue("error", "object.missing", f"{asset.id}: stored object not found at {path}"))
Expand Down
25 changes: 24 additions & 1 deletion visionpack/index/sqlite_index.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,21 @@

from visionpack.core.models import Annotation, Asset, Split


def checkpoint_db(path: Path) -> None:
"""Fold any WAL into ``path`` so the bare file is complete and copyable.

Connections are short-lived so SQLite normally checkpoints on close, but
anything that hashes or copies ``index.db`` as a file (snapshot freeze,
archive pack) must not depend on that — an explicit truncating checkpoint
makes the main file self-contained.
"""
if not path.exists():
return
with closing(sqlite3.connect(path)) as conn:
conn.execute("PRAGMA wal_checkpoint(TRUNCATE)")


_SCHEMA = """
CREATE TABLE IF NOT EXISTS assets (id TEXT PRIMARY KEY, data BLOB NOT NULL);
CREATE TABLE IF NOT EXISTS annotations (id TEXT PRIMARY KEY, asset_id TEXT NOT NULL, data BLOB NOT NULL);
Expand Down Expand Up @@ -85,7 +100,15 @@ def __init__(self, root: Path, db_path: Path | None = None) -> None:

def _connect(self) -> sqlite3.Connection:
self.path.parent.mkdir(parents=True, exist_ok=True)
return sqlite3.connect(self.path)
conn = sqlite3.connect(self.path)
if self._is_live:
# WAL + synchronous=NORMAL is the fast *and* crash-safe combination
# for the live index: commits stop fsyncing the main file on every
# save. Frozen snapshot dbs are left untouched — flipping their
# journal mode would rewrite bytes of a content-addressed file.
conn.execute("PRAGMA journal_mode=WAL")
conn.execute("PRAGMA synchronous=NORMAL")
return conn

def _ensure_schema(self) -> None:
with closing(self._connect()) as conn:
Expand Down
7 changes: 6 additions & 1 deletion visionpack/packing/archive.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from visionpack.core.errors import VisionPackError
from visionpack.core.models import utc_now
from visionpack.core.project import Project
from visionpack.index.sqlite_index import checkpoint_db
from visionpack.stats import collect_stats


Expand Down Expand Up @@ -51,6 +52,7 @@ def pack_archive(project: Project, output: Path | None = None, profile_name: str

index_path = project.root / ".vp" / "db" / "index.db"
if index_path.exists():
checkpoint_db(index_path) # fold any WAL so the packed file is self-contained
files += _add_path(tar, index_path, ".vp/db/index.db")

if include_metadata:
Expand Down Expand Up @@ -115,7 +117,10 @@ def __init__(self, path: Path, compression_level: int) -> None:
def open(self, archive_format: str) -> _TarContext:
self._file = self.path.open("wb")
if archive_format == "tar.zst":
compressor = zstd.ZstdCompressor(level=self.compression_level)
# threads=-1 enables zstd's own multi-threaded compression (one
# worker per core); the archive is a single stream, so this is
# where the parallelism has to come from.
compressor = zstd.ZstdCompressor(level=self.compression_level, threads=-1)
self._zstd = compressor.stream_writer(self._file, closefd=False)
tar = tarfile.open(fileobj=self._zstd, mode="w|")
else:
Expand Down
Loading
Loading