Files
VideoDetect/src/active_learning/label_ingestor.py
T
2026-08-10 19:25:20 -04:00

218 lines
7.4 KiB
Python

"""
Label ingestion for active learning.
Queries the review_queue for annotated labels, extracts face crops,
and builds a versioned training dataset with stratified train/val split.
"""
import csv
import json
import logging
import random
import shutil
from pathlib import Path
from typing import Dict, List, Optional, Tuple
logger = logging.getLogger(__name__)
# Train split fraction (val = 1 - TRAIN_FRAC)
_TRAIN_FRAC = 0.80
class LabelIngestor:
"""Ingest annotated review data and build a versioned dataset."""
def __init__(self, db_connector, config):
self._db = db_connector
self._cfg = config
self._training_path = Path(config.get("storage.training_path", "/data/training"))
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
def count_annotated(self) -> int:
"""Return number of annotated rows ready for training."""
rows = self._db.fetchall(
"SELECT COUNT(*) AS cnt FROM review_queue "
"WHERE annotated = TRUE AND ground_truth IS NOT NULL"
)
return int((rows or [{"cnt": 0}])[0]["cnt"])
def ingest(self, version: str, seed: int = 42) -> Optional[str]:
"""
Build a versioned dataset directory from the review queue.
Returns the dataset path on success, or None if no data is available.
"""
rows = self._fetch_annotated_rows()
if not rows:
logger.warning("No annotated samples found; skipping dataset build")
return None
dataset_path = self._training_path / version
crops_pos = dataset_path / "crops" / "class_1"
crops_neg = dataset_path / "crops" / "class_0"
for d in (crops_pos, crops_neg):
d.mkdir(parents=True, exist_ok=True)
records = self._extract_crops(rows, crops_pos, crops_neg)
if not records:
logger.error("No crops could be extracted from annotated rows")
return None
train_records, val_records = _stratified_split(records, _TRAIN_FRAC, seed)
self._write_labels_csv(dataset_path / "labels.csv", records)
self._write_metadata(dataset_path / "metadata.json", version, records, train_records, val_records)
logger.info(
"Dataset v%s built: %d total (%d train / %d val) at %s",
version, len(records), len(train_records), len(val_records), dataset_path,
)
return str(dataset_path)
# ------------------------------------------------------------------
# Internal helpers
# ------------------------------------------------------------------
def _fetch_annotated_rows(self) -> List[Dict]:
rows = self._db.fetchall(
"""SELECT rq.video_id, rq.ground_truth,
pl.confidence_scores,
v.file_path
FROM review_queue rq
JOIN videos v ON v.id = rq.video_id
LEFT JOIN (
SELECT video_id, confidence_scores,
ROW_NUMBER() OVER (PARTITION BY video_id ORDER BY processed_at DESC) AS rn
FROM processing_logs
) pl ON pl.video_id = rq.video_id AND pl.rn = 1
WHERE rq.annotated = TRUE AND rq.ground_truth IS NOT NULL"""
)
return rows or []
def _extract_crops(
self,
rows: List[Dict],
crops_pos: Path,
crops_neg: Path,
) -> List[Dict]:
"""Copy or re-extract face crops into the dataset directory."""
records: List[Dict] = []
scratch_path = Path(self._cfg.get("storage.scratch_path", "/scratch"))
for row in rows:
video_id = row["video_id"]
label = int(bool(row["ground_truth"]))
dest_dir = crops_pos if label == 1 else crops_neg
try:
contributing = _parse_json_field(row.get("confidence_scores"))
crop_paths = self._resolve_crop_paths(video_id, contributing, scratch_path)
for crop_path in crop_paths:
src = Path(crop_path)
if not src.exists():
logger.debug("Crop not found, skipping: %s", src)
continue
dest = dest_dir / f"video{video_id}_{src.name}"
shutil.copy2(src, dest)
records.append({"crop_path": str(dest), "label": label, "video_id": video_id})
except Exception as exc:
logger.warning("Failed to extract crops for video %s: %s", video_id, exc)
return records
def _resolve_crop_paths(
self,
video_id: int,
contributing_frames: List,
scratch_path: Path,
) -> List[str]:
"""Return paths of existing crop files for this video."""
paths: List[str] = []
# contributing_frames may be a list of frame scores or frame dicts
if contributing_frames:
for item in contributing_frames:
if isinstance(item, dict) and "crop_path" in item:
paths.append(item["crop_path"])
elif isinstance(item, str):
paths.append(item)
# Fall back: glob any crop files already written to scratch
if not paths:
pattern = f"*video{video_id}*.jpg"
paths = [str(p) for p in scratch_path.rglob(pattern)]
return paths
@staticmethod
def _write_labels_csv(csv_path: Path, records: List[Dict]):
with open(csv_path, "w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=["crop_path", "label", "video_id"])
writer.writeheader()
writer.writerows(records)
@staticmethod
def _write_metadata(
meta_path: Path,
version: str,
all_records: List[Dict],
train_records: List[Dict],
val_records: List[Dict],
):
pos = sum(1 for r in all_records if r["label"] == 1)
neg = len(all_records) - pos
meta = {
"version": version,
"total_samples": len(all_records),
"train_samples": len(train_records),
"val_samples": len(val_records),
"class_counts": {"class_0": neg, "class_1": pos},
}
meta_path.write_text(json.dumps(meta, indent=2), encoding="utf-8")
# ------------------------------------------------------------------
# Module-level helpers
# ------------------------------------------------------------------
def _stratified_split(
records: List[Dict],
train_frac: float,
seed: int,
) -> Tuple[List[Dict], List[Dict]]:
"""Return (train, val) with stratification by label."""
rng = random.Random(seed)
by_label: Dict[int, List[Dict]] = {}
for r in records:
by_label.setdefault(r["label"], []).append(r)
train, val = [], []
for label_records in by_label.values():
shuffled = list(label_records)
rng.shuffle(shuffled)
split = max(1, int(len(shuffled) * train_frac))
train.extend(shuffled[:split])
val.extend(shuffled[split:])
rng.shuffle(train)
rng.shuffle(val)
return train, val
def _parse_json_field(value) -> List:
if not value:
return []
if isinstance(value, list):
return value
try:
parsed = json.loads(value)
return parsed if isinstance(parsed, list) else []
except (json.JSONDecodeError, TypeError):
return []