feat: remove object mode pipeline (YOLO, SigLIP, TRAINING_MODE, OBJECT_CLASS)
winnow is a face recognition training tool. Object mode required manual file placement with no Frigate API, pulled in torch/torchvision/transformers/ ultralytics (~2 GB), and was architecturally misaligned with the project goal. Removed: - process_object_mode (YOLO inference), get_yolo_model - SigLIP model stack (get_siglip_model, get_object_embedding, batch variant) - entity_type branching throughout diversity, embeddings, jobs, executor - TRAINING_MODE and OBJECT_CLASS env vars - torch, torchvision, transformers, ultralytics dependencies - pytorch index entries from pyproject.toml - Object mode from README (Modes section, env var table, How It Works)
This commit is contained in:
@@ -39,8 +39,7 @@ Immich library
|
|||||||
│
|
│
|
||||||
▼
|
▼
|
||||||
4. Compute embeddings from the same preview thumbnails
|
4. Compute embeddings from the same preview thumbnails
|
||||||
• Faces → InsightFace (ArcFace / Buffalo_L) → 512-dim vector
|
• InsightFace (ArcFace / Buffalo_L) → 512-dim vector
|
||||||
• Objects → SigLIP (Vision Transformer) → 768-dim vector
|
|
||||||
│
|
│
|
||||||
▼
|
▼
|
||||||
5. Near-duplicate removal — greedy cosine-distance pass drops burst shots
|
5. Near-duplicate removal — greedy cosine-distance pass drops burst shots
|
||||||
@@ -61,35 +60,25 @@ Immich library
|
|||||||
7. Download full-resolution originals from Immich
|
7. Download full-resolution originals from Immich
|
||||||
│
|
│
|
||||||
▼
|
▼
|
||||||
8. Crop and process
|
8. Crop and process — EXIF-corrected, landmark-aligned 112×112 crop (ArcFace format)
|
||||||
• Face mode: EXIF-corrected, landmark-aligned 112×112 crop (ArcFace format)
|
|
||||||
• Object mode: YOLOv9c detection → one crop per matched instance
|
|
||||||
│
|
│
|
||||||
▼
|
▼
|
||||||
9. Deliver
|
9. Deliver — upload crops to Frigate's face registration API
|
||||||
• Face mode: upload crops to Frigate's face registration API
|
↳ below MAX_AUTO_IMAGES — upload, unless the novelty gate
|
||||||
↳ below MAX_AUTO_IMAGES — upload, unless the novelty gate
|
(FRIGATE_SCORE_CEILING) determines the candidate is already
|
||||||
(FRIGATE_SCORE_CEILING) determines the candidate is already
|
covered by the current training set
|
||||||
covered by the current training set
|
↳ at cap + QUALITY_REPLACEMENT=true — with Frigate scoring active,
|
||||||
↳ at cap + QUALITY_REPLACEMENT=true — with Frigate scoring active,
|
swap the most redundant tracked image (highest pre-upload recognize
|
||||||
swap the most redundant tracked image (highest pre-upload recognize
|
score) if the candidate is more novel (lower score); falling back to
|
||||||
score) if the candidate is more novel (lower score); falling back to
|
blur-score comparison when no Frigate scores are available; manually
|
||||||
blur-score comparison when no Frigate scores are available; manually
|
added files are never touched
|
||||||
added files are never touched
|
↳ at cap + QUALITY_REPLACEMENT=false — skip this person
|
||||||
↳ at cap + QUALITY_REPLACEMENT=false — skip this person
|
|
||||||
• Object mode: save crops to disk → place into your Frigate data directory
|
|
||||||
```
|
```
|
||||||
|
|
||||||
Uploaded and rejected asset IDs are persisted across runs. The same image is never processed twice; rejected assets are permanently skipped unless `RETRY_REJECTED=true`.
|
Uploaded and rejected asset IDs are persisted across runs. The same image is never processed twice; rejected assets are permanently skipped unless `RETRY_REJECTED=true`.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Modes
|
|
||||||
|
|
||||||
**Face mode** (default) — extracts face crops using Immich's bounding box metadata, applies EXIF orientation correction, and aligns them to ArcFace's standard 112×112 format using 5-point facial landmarks. Crops are uploaded directly to Frigate's face registration API.
|
|
||||||
|
|
||||||
**Object mode** — runs each full-resolution image through YOLOv9c to detect instances of a target class (dog, cat, car, etc.), crops each detection, and saves it to the output directory. Frigate has no API for uploading object training data; place the crops into your Frigate data directory manually.
|
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Running in Docker
|
## Running in Docker
|
||||||
@@ -116,7 +105,7 @@ services:
|
|||||||
- FRIGATE_URL=http://192.168.1.10:5000
|
- FRIGATE_URL=http://192.168.1.10:5000
|
||||||
- CRON_SCHEDULE=0 3 * * 0
|
- CRON_SCHEDULE=0 3 * * 0
|
||||||
volumes:
|
volumes:
|
||||||
- /path/to/models:/models
|
- /path/to/models:/models # INSIGHTFACE_HOME — persists Buffalo_L model (~300 MB)
|
||||||
- /path/to/cache:/app/.if_cache
|
- /path/to/cache:/app/.if_cache
|
||||||
- /path/to/output:/app/frigate_train
|
- /path/to/output:/app/frigate_train
|
||||||
deploy:
|
deploy:
|
||||||
@@ -162,7 +151,7 @@ See [compose.yml](compose.yml) for the full annotated example with all options.
|
|||||||
| *(empty string)* | Stay alive, run nothing — trigger manually with `docker exec -it winnow winnow` |
|
| *(empty string)* | Stay alive, run nothing — trigger manually with `docker exec -it winnow winnow` |
|
||||||
| Cron expression | Run on startup, then repeat on schedule |
|
| Cron expression | Run on startup, then repeat on schedule |
|
||||||
|
|
||||||
In scheduled mode the process (and loaded models) stays resident between runs. The first run after a fresh install downloads the embedding models (~1–2 GB); subsequent runs use the cached models from the mounted volume.
|
In scheduled mode the process (and loaded models) stays resident between runs. The first run after a fresh install downloads InsightFace Buffalo_L (~300 MB); subsequent runs use the cached model from the mounted volume.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -180,10 +169,8 @@ In scheduled mode the process (and loaded models) stays resident between runs. T
|
|||||||
|
|
||||||
| Variable | Default | Description |
|
| Variable | Default | Description |
|
||||||
| :--- | :--- | :--- |
|
| :--- | :--- | :--- |
|
||||||
| `TRAINING_MODE` | `face` | `face` — upload crops to Frigate; `object` — save crops to disk |
|
|
||||||
| `STRATEGY` | `adaptive` | `adaptive` — embedding-based diversity selection, stops when candidates become redundant; `standard` — fixed 30 images; `broad` — fixed 100 images |
|
| `STRATEGY` | `adaptive` | `adaptive` — embedding-based diversity selection, stops when candidates become redundant; `standard` — fixed 30 images; `broad` — fixed 100 images |
|
||||||
| `LIMIT` | *(unset)* | Exact image count — overrides `STRATEGY` |
|
| `LIMIT` | *(unset)* | Exact image count — overrides `STRATEGY` |
|
||||||
| `OBJECT_CLASS` | `dog` | Target class for object mode (any YOLO class: `dog`, `cat`, `car`, etc.) |
|
|
||||||
| `AUTO_MODE` | *(auto)* | Skip interactive prompts and process all people unattended — auto-detected when no TTY is present (Docker, cron); set `true` to force in a terminal |
|
| `AUTO_MODE` | *(auto)* | Skip interactive prompts and process all people unattended — auto-detected when no TTY is present (Docker, cron); set `true` to force in a terminal |
|
||||||
| `VERBOSE` | `false` | Enable DEBUG-level console output (log file is always DEBUG) |
|
| `VERBOSE` | `false` | Enable DEBUG-level console output (log file is always DEBUG) |
|
||||||
|
|
||||||
@@ -227,7 +214,6 @@ These defaults are tuned for Frigate's ArcFace requirements. winnow will warn on
|
|||||||
| `OPENVINO_DEVICE` | `CPU` | Intel variant only: set `GPU` to use Arc or iGPU; default runs on CPU |
|
| `OPENVINO_DEVICE` | `CPU` | Intel variant only: set `GPU` to use Arc or iGPU; default runs on CPU |
|
||||||
| `ENABLE_CACHE` | `true` | Cache computed embeddings to disk (speeds up re-runs on the same library) |
|
| `ENABLE_CACHE` | `true` | Cache computed embeddings to disk (speeds up re-runs on the same library) |
|
||||||
| `CACHE_DIR` | `.if_cache` | Path for embedding cache and upload tracker files |
|
| `CACHE_DIR` | `.if_cache` | Path for embedding cache and upload tracker files |
|
||||||
| `HF_HOME` | *(system)* | HuggingFace model cache path (SigLIP) |
|
|
||||||
| `INSIGHTFACE_HOME` | *(system)* | InsightFace model cache path (Buffalo_L) |
|
| `INSIGHTFACE_HOME` | *(system)* | InsightFace model cache path (Buffalo_L) |
|
||||||
|
|
||||||
### Output
|
### Output
|
||||||
|
|||||||
+1
-29
@@ -1,7 +1,7 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "winnow"
|
name = "winnow"
|
||||||
version = "0.4.10"
|
version = "0.4.10"
|
||||||
description = "Selects diverse, high-quality photos from Immich as training data for Frigate face recognition and object classification."
|
description = "Selects diverse, high-quality photos from Immich as training data for Frigate face recognition."
|
||||||
license = "AGPL-3.0-or-later"
|
license = "AGPL-3.0-or-later"
|
||||||
requires-python = ">=3.13"
|
requires-python = ">=3.13"
|
||||||
authors = [{ name = "Holden Salomon", email = "holden@arch.fyi" }]
|
authors = [{ name = "Holden Salomon", email = "holden@arch.fyi" }]
|
||||||
@@ -27,10 +27,6 @@ dependencies = [
|
|||||||
"python-dotenv>=1.2.1",
|
"python-dotenv>=1.2.1",
|
||||||
"requests>=2.32.5",
|
"requests>=2.32.5",
|
||||||
"rich>=14.2.0",
|
"rich>=14.2.0",
|
||||||
"torch>=2.12.0",
|
|
||||||
"torchvision>=0.27.0",
|
|
||||||
"transformers>=5.12.0",
|
|
||||||
"ultralytics>=8.4.66",
|
|
||||||
]
|
]
|
||||||
|
|
||||||
[project.scripts]
|
[project.scripts]
|
||||||
@@ -53,27 +49,6 @@ required-environments = [
|
|||||||
"sys_platform == 'linux' and platform_machine == 'aarch64'",
|
"sys_platform == 'linux' and platform_machine == 'aarch64'",
|
||||||
]
|
]
|
||||||
|
|
||||||
[tool.uv.sources]
|
|
||||||
torch = [
|
|
||||||
{ index = "pytorch-cu126", marker = "sys_platform == 'linux' and platform_machine == 'x86_64'" },
|
|
||||||
{ index = "pytorch-cpu", marker = "sys_platform == 'linux' and platform_machine == 'aarch64'" },
|
|
||||||
{ index = "pytorch-cpu", marker = "sys_platform != 'linux'" },
|
|
||||||
]
|
|
||||||
torchvision = [
|
|
||||||
{ index = "pytorch-cu126", marker = "sys_platform == 'linux' and platform_machine == 'x86_64'" },
|
|
||||||
{ index = "pytorch-cpu", marker = "sys_platform == 'linux' and platform_machine == 'aarch64'" },
|
|
||||||
{ index = "pytorch-cpu", marker = "sys_platform != 'linux'" },
|
|
||||||
]
|
|
||||||
|
|
||||||
[[tool.uv.index]]
|
|
||||||
name = "pytorch-cu126"
|
|
||||||
url = "https://download.pytorch.org/whl/cu126"
|
|
||||||
explicit = true
|
|
||||||
|
|
||||||
[[tool.uv.index]]
|
|
||||||
name = "pytorch-cpu"
|
|
||||||
url = "https://download.pytorch.org/whl/cpu"
|
|
||||||
explicit = true
|
|
||||||
|
|
||||||
[dependency-groups]
|
[dependency-groups]
|
||||||
dev = [
|
dev = [
|
||||||
@@ -104,9 +79,6 @@ nvidia-cudnn-cu12 = "nvidia.cudnn"
|
|||||||
onnxruntime-gpu = "onnxruntime"
|
onnxruntime-gpu = "onnxruntime"
|
||||||
requests = "requests"
|
requests = "requests"
|
||||||
rich = "rich"
|
rich = "rich"
|
||||||
torch = "torch"
|
|
||||||
transformers = "transformers"
|
|
||||||
ultralytics = "ultralytics"
|
|
||||||
|
|
||||||
[tool.deptry.per_rule_ignores]
|
[tool.deptry.per_rule_ignores]
|
||||||
DEP002 = ["onnxruntime-gpu", "nvidia-cudnn-cu12"]
|
DEP002 = ["onnxruntime-gpu", "nvidia-cudnn-cu12"]
|
||||||
|
|||||||
+32
-55
@@ -5,7 +5,7 @@ Selection pipeline:
|
|||||||
1. Concurrent thumbnail download
|
1. Concurrent thumbnail download
|
||||||
2. Quality filtering (blur, IR, exposure, confidence, face size)
|
2. Quality filtering (blur, IR, exposure, confidence, face size)
|
||||||
3. Face crop extraction (embed person's face, not full image)
|
3. Face crop extraction (embed person's face, not full image)
|
||||||
4. Embedding computation (InsightFace or SigLIP)
|
4. Embedding computation (InsightFace)
|
||||||
5. Cluster-aware selection (K-Medoids + FPS with hard example weighting)
|
5. Cluster-aware selection (K-Medoids + FPS with hard example weighting)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -28,7 +28,6 @@ def select_diverse_assets(
|
|||||||
limit: int | str,
|
limit: int | str,
|
||||||
entity_name: str,
|
entity_name: str,
|
||||||
selection_mode: str = "smart",
|
selection_mode: str = "smart",
|
||||||
entity_type: str = "face",
|
|
||||||
person_id: str | None = None,
|
person_id: str | None = None,
|
||||||
progress_callback=None,
|
progress_callback=None,
|
||||||
) -> list:
|
) -> list:
|
||||||
@@ -38,9 +37,8 @@ def select_diverse_assets(
|
|||||||
Args:
|
Args:
|
||||||
assets: List of asset dicts from Immich API
|
assets: List of asset dicts from Immich API
|
||||||
limit: Number to select, or "auto" for dynamic selection
|
limit: Number to select, or "auto" for dynamic selection
|
||||||
entity_name: Name of the person/object for logging
|
entity_name: Name of the person for logging
|
||||||
selection_mode: 'smart' (embedding-based) or 'time' (time spread)
|
selection_mode: 'smart' (embedding-based) or 'time' (time spread)
|
||||||
entity_type: 'face' or 'object' - determines embedding model
|
|
||||||
progress_callback: Optional callback(current, total) for progress
|
progress_callback: Optional callback(current, total) for progress
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -53,14 +51,13 @@ def select_diverse_assets(
|
|||||||
# Sort by creation time
|
# Sort by creation time
|
||||||
assets = sorted(assets, key=lambda x: x.get("fileCreatedAt", ""))
|
assets = sorted(assets, key=lambda x: x.get("fileCreatedAt", ""))
|
||||||
|
|
||||||
if selection_mode != "smart" or not is_embedding_available(entity_type):
|
if selection_mode != "smart" or not is_embedding_available():
|
||||||
if selection_mode == "smart":
|
if selection_mode == "smart":
|
||||||
model_name = "InsightFace" if entity_type == "face" else "SigLIP"
|
logger.warning("InsightFace unavailable. Falling back to time spread.")
|
||||||
logger.warning(f"{model_name} unavailable. Falling back to time spread.")
|
|
||||||
return _select_time_spread(assets, limit)
|
return _select_time_spread(assets, limit)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
return _select_by_embedding(assets, limit, entity_type, person_id, progress_callback)
|
return _select_by_embedding(assets, limit, person_id, progress_callback)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Smart Diversity failed: {e}. Falling back to time spread.")
|
logger.error(f"Smart Diversity failed: {e}. Falling back to time spread.")
|
||||||
return _select_time_spread(assets, limit)
|
return _select_time_spread(assets, limit)
|
||||||
@@ -179,7 +176,6 @@ def _crop_face_from_thumbnail(
|
|||||||
def _select_by_embedding(
|
def _select_by_embedding(
|
||||||
assets: list,
|
assets: list,
|
||||||
limit: int | str,
|
limit: int | str,
|
||||||
entity_type: str,
|
|
||||||
person_id: str | None = None,
|
person_id: str | None = None,
|
||||||
progress_callback=None,
|
progress_callback=None,
|
||||||
) -> list:
|
) -> list:
|
||||||
@@ -188,7 +184,7 @@ def _select_by_embedding(
|
|||||||
Pipeline:
|
Pipeline:
|
||||||
1. Concurrent thumbnail download
|
1. Concurrent thumbnail download
|
||||||
2. Quality filtering
|
2. Quality filtering
|
||||||
3. Face crop extraction (face mode only)
|
3. Face crop extraction
|
||||||
4. Embedding computation
|
4. Embedding computation
|
||||||
5. Cluster-aware selection with hard example weighting
|
5. Cluster-aware selection with hard example weighting
|
||||||
"""
|
"""
|
||||||
@@ -243,28 +239,25 @@ def _select_by_embedding(
|
|||||||
|
|
||||||
confidence = _get_face_confidence(asset, person_id=person_id)
|
confidence = _get_face_confidence(asset, person_id=person_id)
|
||||||
|
|
||||||
if entity_type == "face":
|
face_bbox = _get_face_bbox(asset, person_id=person_id)
|
||||||
face_bbox = _get_face_bbox(asset, person_id=person_id)
|
quality = assess_quality(
|
||||||
quality = assess_quality(
|
img,
|
||||||
img,
|
face_bbox=face_bbox,
|
||||||
face_bbox=face_bbox,
|
confidence=confidence,
|
||||||
confidence=confidence,
|
blur_threshold=Config.BLUR_THRESHOLD,
|
||||||
blur_threshold=Config.BLUR_THRESHOLD,
|
min_face_px=Config.MIN_FACE_WIDTH,
|
||||||
min_face_px=Config.MIN_FACE_WIDTH,
|
min_confidence=Config.MIN_CONFIDENCE,
|
||||||
min_confidence=Config.MIN_CONFIDENCE,
|
)
|
||||||
)
|
if not quality.passed:
|
||||||
if not quality.passed:
|
quality_filtered += 1
|
||||||
quality_filtered += 1
|
logger.debug(f"Quality filtered {asset['id']}: {quality.reason}")
|
||||||
logger.debug(f"Quality filtered {asset['id']}: {quality.reason}")
|
continue
|
||||||
continue
|
|
||||||
|
|
||||||
asset["quality_score"] = quality.blur_score
|
asset["quality_score"] = quality.blur_score
|
||||||
face_crop = _crop_face_from_thumbnail(img, asset, person_id=person_id)
|
face_crop = _crop_face_from_thumbnail(img, asset, person_id=person_id)
|
||||||
embed_img = face_crop if face_crop is not None else img
|
embed_img = face_crop if face_crop is not None else img
|
||||||
else:
|
|
||||||
embed_img = img
|
|
||||||
|
|
||||||
emb = get_embedding(embed_img, entity_type, asset_id=asset["id"])
|
emb = get_embedding(embed_img, asset_id=asset["id"])
|
||||||
if emb is not None:
|
if emb is not None:
|
||||||
embeddings.append(emb)
|
embeddings.append(emb)
|
||||||
valid_candidates.append(asset)
|
valid_candidates.append(asset)
|
||||||
@@ -300,7 +293,6 @@ def _select_by_embedding(
|
|||||||
embeddings,
|
embeddings,
|
||||||
valid_candidates,
|
valid_candidates,
|
||||||
limit,
|
limit,
|
||||||
entity_type=entity_type,
|
|
||||||
confidence_scores=confidence_scores,
|
confidence_scores=confidence_scores,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -429,11 +421,11 @@ def _kmedoids(dist_matrix: np.ndarray, k: int, max_iter: int = 50) -> tuple[list
|
|||||||
# =============================================================================
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
def _compute_adaptive_threshold(emb_normed: np.ndarray, entity_type: str) -> float:
|
def _compute_adaptive_threshold(emb_normed: np.ndarray) -> float:
|
||||||
"""Compute adaptive FPS stop threshold based on actual embedding distribution.
|
"""Compute adaptive FPS stop threshold based on actual embedding distribution.
|
||||||
|
|
||||||
Instead of a hardcoded threshold, samples pairwise distances and sets
|
Instead of a hardcoded threshold, samples pairwise distances and sets
|
||||||
the threshold as a fraction of the median pairwise distance.
|
the threshold as 20% of the median pairwise distance.
|
||||||
"""
|
"""
|
||||||
n = len(emb_normed)
|
n = len(emb_normed)
|
||||||
sample_size = min(200, n)
|
sample_size = min(200, n)
|
||||||
@@ -441,22 +433,14 @@ def _compute_adaptive_threshold(emb_normed: np.ndarray, entity_type: str) -> flo
|
|||||||
indices = rng.choice(n, sample_size, replace=False) if n > sample_size else np.arange(n)
|
indices = rng.choice(n, sample_size, replace=False) if n > sample_size else np.arange(n)
|
||||||
sample = emb_normed[indices]
|
sample = emb_normed[indices]
|
||||||
|
|
||||||
# Compute pairwise cosine distances for the sample
|
|
||||||
pairwise = 1 - sample @ sample.T
|
pairwise = 1 - sample @ sample.T
|
||||||
upper_tri = pairwise[np.triu_indices(len(sample), k=1)]
|
upper_tri = pairwise[np.triu_indices(len(sample), k=1)]
|
||||||
if len(upper_tri) == 0:
|
if len(upper_tri) == 0:
|
||||||
return 0.05
|
return 0.05
|
||||||
median_dist = float(np.median(upper_tri))
|
median_dist = float(np.median(upper_tri))
|
||||||
|
threshold = max(0.05, median_dist * 0.20)
|
||||||
|
|
||||||
# Faces: 20% of median (tighter — want fewer, more distinct images)
|
logger.debug(f"Adaptive threshold: {threshold:.4f} (median_dist={median_dist:.4f})")
|
||||||
# Objects: 10% of median (wider — want more diversity)
|
|
||||||
fraction = 0.20 if entity_type == "face" else 0.10
|
|
||||||
threshold = max(0.05, median_dist * fraction)
|
|
||||||
|
|
||||||
logger.debug(
|
|
||||||
f"Adaptive threshold: {threshold:.4f} "
|
|
||||||
f"(median_dist={median_dist:.4f}, fraction={fraction}, type={entity_type})"
|
|
||||||
)
|
|
||||||
return threshold
|
return threshold
|
||||||
|
|
||||||
|
|
||||||
@@ -464,7 +448,6 @@ def _cluster_aware_selection(
|
|||||||
embeddings: list,
|
embeddings: list,
|
||||||
candidates: list,
|
candidates: list,
|
||||||
limit: int | str,
|
limit: int | str,
|
||||||
entity_type: str = "face",
|
|
||||||
confidence_scores: list | None = None,
|
confidence_scores: list | None = None,
|
||||||
) -> list:
|
) -> list:
|
||||||
"""Two-stage selection: K-Medoids clustering → FPS with hard example weighting.
|
"""Two-stage selection: K-Medoids clustering → FPS with hard example weighting.
|
||||||
@@ -484,13 +467,13 @@ def _cluster_aware_selection(
|
|||||||
|
|
||||||
# Build confidence weight array for hard example boosting
|
# Build confidence weight array for hard example boosting
|
||||||
conf_array = np.ones(n)
|
conf_array = np.ones(n)
|
||||||
if confidence_scores and entity_type == "face":
|
if confidence_scores:
|
||||||
for i, c in enumerate(confidence_scores):
|
for i, c in enumerate(confidence_scores):
|
||||||
if c is not None:
|
if c is not None:
|
||||||
conf_array[i] = c
|
conf_array[i] = c
|
||||||
|
|
||||||
# Compute adaptive threshold for auto mode
|
# Compute adaptive threshold for auto mode
|
||||||
auto_threshold = _compute_adaptive_threshold(emb_normed, entity_type) if limit == "auto" else 0.0
|
auto_threshold = _compute_adaptive_threshold(emb_normed) if limit == "auto" else 0.0
|
||||||
target = Config.MAX_AUTO_IMAGES if limit == "auto" else limit
|
target = Config.MAX_AUTO_IMAGES if limit == "auto" else limit
|
||||||
|
|
||||||
# --- Stage 1: K-Medoids clustering ---
|
# --- Stage 1: K-Medoids clustering ---
|
||||||
@@ -542,15 +525,9 @@ def _cluster_aware_selection(
|
|||||||
min_dists = np.minimum(min_dists, dists_to_new)
|
min_dists = np.minimum(min_dists, dists_to_new)
|
||||||
min_dists[best_idx] = -np.inf
|
min_dists[best_idx] = -np.inf
|
||||||
|
|
||||||
# Log hard example stats
|
selected_conf = [conf_array[i] for i in selected if conf_array[i] < 1.0]
|
||||||
if entity_type == "face":
|
hard_count = sum(1 for c in selected_conf if c < 0.85)
|
||||||
selected_conf = [conf_array[i] for i in selected if conf_array[i] < 1.0]
|
logger.info(f"Selection complete: {len(selected)} images ({hard_count} hard examples with confidence < 0.85).")
|
||||||
hard_count = sum(1 for c in selected_conf if c < 0.85)
|
|
||||||
logger.info(
|
|
||||||
f"Selection complete: {len(selected)} images " f"({hard_count} hard examples with confidence < 0.85)."
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
logger.info(f"Selection complete: {len(selected)} diverse images.")
|
|
||||||
|
|
||||||
return [candidates[i] for i in selected]
|
return [candidates[i] for i in selected]
|
||||||
|
|
||||||
|
|||||||
+18
-168
@@ -1,8 +1,7 @@
|
|||||||
"""
|
"""
|
||||||
Unified embedding interface for faces and objects.
|
Embedding interface for face diversity selection.
|
||||||
|
|
||||||
- Faces: InsightFace (ArcFace/Buffalo_L) — or reuse from Immich
|
- Faces: InsightFace (ArcFace/Buffalo_L) — or reuse from Immich
|
||||||
- Objects: SigLIP (Vision Transformer via transformers)
|
|
||||||
- Caching: Disk-based cache avoids recomputation on reruns
|
- Caching: Disk-based cache avoids recomputation on reruns
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -42,12 +41,9 @@ def _suppress_output():
|
|||||||
os.close(saved_err)
|
os.close(saved_err)
|
||||||
|
|
||||||
|
|
||||||
# Lazy-loaded singletons
|
# Lazy-loaded singleton
|
||||||
_insightface_app = None
|
_insightface_app = None
|
||||||
_insightface_loaded = False
|
_insightface_loaded = False
|
||||||
_siglip_model = None
|
|
||||||
_siglip_processor = None
|
|
||||||
_siglip_loaded = False
|
|
||||||
|
|
||||||
|
|
||||||
def _is_force_cpu() -> bool:
|
def _is_force_cpu() -> bool:
|
||||||
@@ -202,169 +198,41 @@ def get_face_embedding(img_pil: Image.Image) -> np.ndarray | None:
|
|||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
# SigLIP (Objects)
|
# Embedding Interface with Caching
|
||||||
# =============================================================================
|
|
||||||
|
|
||||||
|
|
||||||
def get_siglip_model():
|
|
||||||
"""Singleton for SigLIP model and processor with GPU auto-detection."""
|
|
||||||
global _siglip_model, _siglip_processor, _siglip_loaded
|
|
||||||
if _siglip_loaded:
|
|
||||||
return _siglip_model, _siglip_processor
|
|
||||||
_siglip_loaded = True
|
|
||||||
|
|
||||||
try:
|
|
||||||
import warnings
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from transformers import AutoImageProcessor, SiglipVisionModel
|
|
||||||
|
|
||||||
model_name = "google/siglip-base-patch16-224"
|
|
||||||
|
|
||||||
# Disk cache check — path derived from model_name using HuggingFace's slug convention
|
|
||||||
hf_home = os.environ.get("HF_HOME", os.path.join(os.path.expanduser("~"), ".cache", "huggingface"))
|
|
||||||
cache_slug = "models--" + model_name.replace("/", "--")
|
|
||||||
model_cache = Path(hf_home) / "hub" / cache_slug
|
|
||||||
if model_cache.exists() and any(model_cache.iterdir()):
|
|
||||||
logger.info(f"SigLIP {model_name}: found in model cache")
|
|
||||||
else:
|
|
||||||
logger.info(f"SigLIP {model_name}: not cached — downloading now (~380 MB)")
|
|
||||||
|
|
||||||
logger.info(f"SigLIP {model_name}: loading into memory...")
|
|
||||||
t0 = time.time()
|
|
||||||
|
|
||||||
with warnings.catch_warnings():
|
|
||||||
warnings.filterwarnings("ignore", category=FutureWarning)
|
|
||||||
warnings.filterwarnings("ignore", message=".*use_fast.*")
|
|
||||||
_siglip_processor = AutoImageProcessor.from_pretrained(model_name, use_fast=True)
|
|
||||||
_siglip_model = SiglipVisionModel.from_pretrained(model_name)
|
|
||||||
|
|
||||||
_siglip_model.eval()
|
|
||||||
|
|
||||||
# Move to GPU if available (ROCm builds expose torch.cuda.is_available() == True)
|
|
||||||
if not _is_force_cpu():
|
|
||||||
if torch.cuda.is_available():
|
|
||||||
_siglip_model = _siglip_model.cuda()
|
|
||||||
device_name = "CUDA GPU"
|
|
||||||
elif hasattr(torch, "xpu") and torch.xpu.is_available():
|
|
||||||
_siglip_model = _siglip_model.to("xpu")
|
|
||||||
device_name = "Intel XPU"
|
|
||||||
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
|
||||||
_siglip_model = _siglip_model.to("mps")
|
|
||||||
device_name = "Apple MPS"
|
|
||||||
else:
|
|
||||||
device_name = "CPU"
|
|
||||||
else:
|
|
||||||
device_name = "CPU (FORCE_CPU)"
|
|
||||||
|
|
||||||
logger.info(f"SigLIP {model_name}: ready on {device_name} ({time.time() - t0:.1f}s)")
|
|
||||||
return _siglip_model, _siglip_processor
|
|
||||||
|
|
||||||
except ImportError as e:
|
|
||||||
logger.error(f"transformers/torch not installed: {e}")
|
|
||||||
return None, None
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Failed to load SigLIP: {e}")
|
|
||||||
return None, None
|
|
||||||
|
|
||||||
|
|
||||||
def get_object_embedding(img_pil: Image.Image) -> np.ndarray | None:
|
|
||||||
"""Get 768-dim SigLIP embedding for an image."""
|
|
||||||
model, processor = get_siglip_model()
|
|
||||||
if model is None:
|
|
||||||
return None
|
|
||||||
|
|
||||||
try:
|
|
||||||
import torch
|
|
||||||
|
|
||||||
inputs = processor(images=img_pil, return_tensors="pt")
|
|
||||||
device = next(model.parameters()).device
|
|
||||||
inputs = {k: v.to(device) for k, v in inputs.items()}
|
|
||||||
|
|
||||||
with torch.no_grad():
|
|
||||||
outputs = model(**inputs)
|
|
||||||
return outputs.pooler_output.squeeze().cpu().numpy()
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error getting object embedding: {e}")
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def get_object_embeddings_batch(images: list[Image.Image]) -> list[np.ndarray | None]:
|
|
||||||
"""Get SigLIP embeddings for a batch of images (GPU-efficient)."""
|
|
||||||
model, processor = get_siglip_model()
|
|
||||||
if model is None:
|
|
||||||
return [None] * len(images)
|
|
||||||
|
|
||||||
try:
|
|
||||||
import torch
|
|
||||||
|
|
||||||
inputs = processor(images=images, return_tensors="pt", padding=True)
|
|
||||||
device = next(model.parameters()).device
|
|
||||||
inputs = {k: v.to(device) for k, v in inputs.items()}
|
|
||||||
|
|
||||||
with torch.no_grad():
|
|
||||||
outputs = model(**inputs)
|
|
||||||
embeddings = outputs.pooler_output.cpu().numpy()
|
|
||||||
return [embeddings[i] for i in range(len(embeddings))]
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error in batch embedding: {e}")
|
|
||||||
# Fall back to individual computation
|
|
||||||
return [get_object_embedding(img) for img in images]
|
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
|
||||||
# Unified Interface with Caching
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
def get_embedding(
|
def get_embedding(
|
||||||
img_pil: Image.Image,
|
img_pil: Image.Image,
|
||||||
entity_type: str = "face",
|
|
||||||
asset_id: str | None = None,
|
asset_id: str | None = None,
|
||||||
immich_embedding: np.ndarray | None = None,
|
immich_embedding: np.ndarray | None = None,
|
||||||
) -> np.ndarray | None:
|
) -> np.ndarray | None:
|
||||||
"""Get embedding for an image based on entity type.
|
"""Get embedding for a face image.
|
||||||
|
|
||||||
Priority:
|
Priority:
|
||||||
1. Pre-fetched Immich embedding (if provided)
|
1. Pre-fetched Immich embedding (if provided)
|
||||||
2. Disk cache (if enabled and asset_id provided)
|
2. Disk cache (if enabled and asset_id provided)
|
||||||
3. Local model computation (InsightFace or SigLIP)
|
3. Local InsightFace computation
|
||||||
|
|
||||||
Args:
|
|
||||||
img_pil: The image to embed
|
|
||||||
entity_type: 'face' or 'object'
|
|
||||||
asset_id: Optional asset ID for cache lookup
|
|
||||||
immich_embedding: Optional pre-fetched embedding from Immich API
|
|
||||||
"""
|
"""
|
||||||
from .config import Config
|
from .config import Config
|
||||||
|
|
||||||
use_cache = Config.ENABLE_CACHE and asset_id is not None
|
use_cache = Config.ENABLE_CACHE and asset_id is not None
|
||||||
cache = get_cache(Config.CACHE_DIR) if use_cache else None
|
cache = get_cache(Config.CACHE_DIR) if use_cache else None
|
||||||
# Use a single consistent cache key per model so lookups and stores always match.
|
|
||||||
# "immich" was previously used as the face key on the lookup path but "insightface"
|
|
||||||
# on the store path — meaning the cache was never hit for locally-computed embeddings.
|
|
||||||
cache_key = "insightface" if entity_type == "face" else "siglip"
|
|
||||||
|
|
||||||
# 1. Use Immich embedding if provided
|
|
||||||
if immich_embedding is not None:
|
if immich_embedding is not None:
|
||||||
if cache:
|
if cache:
|
||||||
cache.put(asset_id, immich_embedding, cache_key)
|
cache.put(asset_id, immich_embedding, "insightface")
|
||||||
return immich_embedding
|
return immich_embedding
|
||||||
|
|
||||||
# 2. Check disk cache
|
|
||||||
if cache:
|
if cache:
|
||||||
cached = cache.get(asset_id, cache_key)
|
cached = cache.get(asset_id, "insightface")
|
||||||
if cached is not None:
|
if cached is not None:
|
||||||
return cached
|
return cached
|
||||||
|
|
||||||
# 3. Compute locally
|
emb = get_face_embedding(img_pil)
|
||||||
if entity_type == "face":
|
|
||||||
emb = get_face_embedding(img_pil)
|
|
||||||
else:
|
|
||||||
emb = get_object_embedding(img_pil)
|
|
||||||
|
|
||||||
if emb is not None and cache:
|
if emb is not None and cache:
|
||||||
cache.put(asset_id, emb, cache_key)
|
cache.put(asset_id, emb, "insightface")
|
||||||
|
|
||||||
return emb
|
return emb
|
||||||
|
|
||||||
@@ -378,36 +246,18 @@ def _is_module_available(module_name: str) -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def is_embedding_available(entity_type: str = "face", *, load: bool = False) -> bool:
|
def is_embedding_available(*, load: bool = False) -> bool:
|
||||||
"""Check if embedding model is available for the given entity type.
|
"""Check if InsightFace is available.
|
||||||
|
|
||||||
By default this performs a lightweight import-check only (no model loading).
|
By default this performs a lightweight import-check only (no model loading).
|
||||||
Pass ``load=True`` to actually load the model (expensive, hundreds of MB).
|
Pass ``load=True`` to actually load the model (expensive, ~300 MB).
|
||||||
|
|
||||||
Args:
|
|
||||||
entity_type: 'face' or 'object'
|
|
||||||
load: If True, fully load the model to verify. If False (default),
|
|
||||||
only check that the required packages are importable.
|
|
||||||
"""
|
"""
|
||||||
if load:
|
if load:
|
||||||
if entity_type == "face":
|
|
||||||
return get_insightface_app() is not None
|
|
||||||
model, _ = get_siglip_model()
|
|
||||||
return model is not None
|
|
||||||
|
|
||||||
# Lightweight check: just verify the packages are importable
|
|
||||||
if entity_type == "face":
|
|
||||||
return _is_module_available("insightface") and _is_module_available("onnxruntime")
|
|
||||||
return _is_module_available("transformers") and _is_module_available("torch")
|
|
||||||
|
|
||||||
|
|
||||||
def load_embedding_model(entity_type: str = "face") -> bool:
|
|
||||||
"""Explicitly load the embedding model for the given entity type.
|
|
||||||
|
|
||||||
Returns True if the model loaded successfully.
|
|
||||||
"""
|
|
||||||
if entity_type == "face":
|
|
||||||
return get_insightface_app() is not None
|
return get_insightface_app() is not None
|
||||||
model, _ = get_siglip_model()
|
return _is_module_available("insightface") and _is_module_available("onnxruntime")
|
||||||
return model is not None
|
|
||||||
|
|
||||||
|
def load_embedding_model() -> bool:
|
||||||
|
"""Explicitly load InsightFace. Returns True if the model loaded successfully."""
|
||||||
|
return get_insightface_app() is not None
|
||||||
|
|
||||||
|
|||||||
+10
-37
@@ -19,7 +19,7 @@ from .frigate_api import (
|
|||||||
get_frigate_person_files,
|
get_frigate_person_files,
|
||||||
recognize_face,
|
recognize_face,
|
||||||
)
|
)
|
||||||
from .image_processing import process_face_mode, process_full_mode, process_object_mode
|
from .image_processing import process_face_mode
|
||||||
from .immich_api import fetch_face_data, fetch_full_image
|
from .immich_api import fetch_face_data, fetch_full_image
|
||||||
from .log_config import console
|
from .log_config import console
|
||||||
from .quality import assess_quality
|
from .quality import assess_quality
|
||||||
@@ -250,26 +250,21 @@ def execute_jobs(jobs: list[dict]) -> None:
|
|||||||
if img is None:
|
if img is None:
|
||||||
progress.console.print(f"[red]Failed download {asset['id']}[/red]")
|
progress.console.print(f"[red]Failed download {asset['id']}[/red]")
|
||||||
else:
|
else:
|
||||||
saved = (
|
saved = process_face_mode(
|
||||||
process_face_mode(img, asset, person, person_dir, count, insightface_app=insightface_app)
|
img, asset, person, person_dir, count, insightface_app=insightface_app
|
||||||
if mode == "face"
|
|
||||||
else process_object_mode(img, config, person_dir, count)
|
|
||||||
if mode == "object"
|
|
||||||
else process_full_mode(img, person_dir, count)
|
|
||||||
)
|
)
|
||||||
if saved:
|
if saved:
|
||||||
# Record which asset produced which output file
|
|
||||||
filename = f"{count}.jpg"
|
filename = f"{count}.jpg"
|
||||||
asset_map[filename] = asset["id"]
|
asset_map[filename] = asset["id"]
|
||||||
score_map[filename] = asset.get("quality_score")
|
score_map[filename] = asset.get("quality_score")
|
||||||
if mode == "face" and isinstance(saved, tuple):
|
if isinstance(saved, tuple):
|
||||||
dims_map[filename] = saved
|
dims_map[filename] = saved
|
||||||
# Time-spread path: compute blur score from the downloaded
|
# Time-spread path: compute blur score from the downloaded
|
||||||
# image. Cap at 1440px so the scale matches the preview
|
# image. Cap at 1440px so the scale matches the preview
|
||||||
# thumbnails the embedding path uses for scoring — Laplacian
|
# thumbnails the embedding path uses for scoring — Laplacian
|
||||||
# variance grows with resolution, making full-res and
|
# variance grows with resolution, making full-res and
|
||||||
# thumbnail scores incomparable if left uncapped.
|
# thumbnail scores incomparable if left uncapped.
|
||||||
if mode == "face" and score_map[filename] is None:
|
if score_map[filename] is None:
|
||||||
try:
|
try:
|
||||||
score_img = img.convert("RGB") if img.mode != "RGB" else img
|
score_img = img.convert("RGB") if img.mode != "RGB" else img
|
||||||
if score_img.width > 1440 or score_img.height > 1440:
|
if score_img.width > 1440 or score_img.height > 1440:
|
||||||
@@ -279,12 +274,6 @@ def execute_jobs(jobs: list[dict]) -> None:
|
|||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.debug(f"Quality score fallback for {asset['id']}: {exc}")
|
logger.debug(f"Quality score fallback for {asset['id']}: {exc}")
|
||||||
score_map[filename] = 0.0 # unknown quality — treat as lowest
|
score_map[filename] = 0.0 # unknown quality — treat as lowest
|
||||||
# Also record object-mode variant filenames
|
|
||||||
if mode == "object":
|
|
||||||
for f in sorted(os.listdir(person_dir)):
|
|
||||||
if f.startswith(f"{count}_") and f not in asset_map:
|
|
||||||
asset_map[f] = asset["id"]
|
|
||||||
score_map[f] = asset.get("face_confidence")
|
|
||||||
|
|
||||||
count += 1
|
count += 1
|
||||||
else:
|
else:
|
||||||
@@ -312,29 +301,13 @@ def execute_jobs(jobs: list[dict]) -> None:
|
|||||||
def upload_to_frigate(jobs: list[dict]) -> None:
|
def upload_to_frigate(jobs: list[dict]) -> None:
|
||||||
"""Upload processed face crops to Frigate via API with detailed logging.
|
"""Upload processed face crops to Frigate via API with detailed logging.
|
||||||
|
|
||||||
Only runs for face-mode jobs. Object-mode crops are saved to the output
|
|
||||||
directory as the deliverable and must be copied to Frigate manually.
|
|
||||||
|
|
||||||
After each successful upload, records the Immich asset ID in the
|
After each successful upload, records the Immich asset ID in the
|
||||||
upload tracker so it is skipped on future runs.
|
upload tracker so it is skipped on future runs.
|
||||||
"""
|
"""
|
||||||
face_jobs = [j for j in jobs if j["config"].get("mode", "face") == "face"]
|
if not jobs:
|
||||||
|
rprint("[dim]No jobs to upload.[/dim]")
|
||||||
if not face_jobs:
|
|
||||||
rprint("[dim]No face-mode jobs to upload.[/dim]")
|
|
||||||
return
|
return
|
||||||
|
|
||||||
# Notify user about object-mode jobs that were skipped
|
|
||||||
object_jobs = [j for j in jobs if j["config"].get("mode") == "object"]
|
|
||||||
for job in object_jobs:
|
|
||||||
name = job["person"]["name"]
|
|
||||||
try:
|
|
||||||
person_dir = _safe_person_dir(Config.OUTPUT_DIR, name)
|
|
||||||
except ValueError as e:
|
|
||||||
logger.error(str(e))
|
|
||||||
continue
|
|
||||||
rprint(f" [dim]📁 {name} (object): crops saved to {person_dir} — copy to Frigate manually[/dim]")
|
|
||||||
|
|
||||||
frigate_url = os.environ.get("FRIGATE_URL", "")
|
frigate_url = os.environ.get("FRIGATE_URL", "")
|
||||||
if not frigate_url:
|
if not frigate_url:
|
||||||
rprint("[yellow]⚠️ FRIGATE_URL not set, skipping upload.[/yellow]")
|
rprint("[yellow]⚠️ FRIGATE_URL not set, skipping upload.[/yellow]")
|
||||||
@@ -347,7 +320,7 @@ def upload_to_frigate(jobs: list[dict]) -> None:
|
|||||||
# from the asset_map stored on each job during execute_jobs()
|
# from the asset_map stored on each job during execute_jobs()
|
||||||
filename_to_asset_id: dict[str, dict[str, str]] = {}
|
filename_to_asset_id: dict[str, dict[str, str]] = {}
|
||||||
total_files = 0
|
total_files = 0
|
||||||
for job in face_jobs:
|
for job in jobs:
|
||||||
name = job["person"]["name"]
|
name = job["person"]["name"]
|
||||||
asset_map = job.get("asset_map", {})
|
asset_map = job.get("asset_map", {})
|
||||||
filename_to_asset_id[name] = asset_map
|
filename_to_asset_id[name] = asset_map
|
||||||
@@ -357,7 +330,7 @@ def upload_to_frigate(jobs: list[dict]) -> None:
|
|||||||
rprint(" [yellow]No images found to upload.[/yellow]")
|
rprint(" [yellow]No images found to upload.[/yellow]")
|
||||||
return
|
return
|
||||||
|
|
||||||
rprint(f" People: [bold]{len(face_jobs)}[/bold], Total images: [bold]{total_files}[/bold]")
|
rprint(f" People: [bold]{len(jobs)}[/bold], Total images: [bold]{total_files}[/bold]")
|
||||||
|
|
||||||
uploaded, failed = 0, 0
|
uploaded, failed = 0, 0
|
||||||
max_retries = 2
|
max_retries = 2
|
||||||
@@ -375,7 +348,7 @@ def upload_to_frigate(jobs: list[dict]) -> None:
|
|||||||
) as progress:
|
) as progress:
|
||||||
upload_task = progress.add_task("[green]Uploading to Frigate", total=total_files)
|
upload_task = progress.add_task("[green]Uploading to Frigate", total=total_files)
|
||||||
|
|
||||||
for job in face_jobs:
|
for job in jobs:
|
||||||
name = job["person"]["name"]
|
name = job["person"]["name"]
|
||||||
# URL-encode the name for the API (handles spaces, special chars)
|
# URL-encode the name for the API (handles spaces, special chars)
|
||||||
encoded_name = quote(name, safe="")
|
encoded_name = quote(name, safe="")
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
"""Image processing functions for cropping faces and objects."""
|
"""Image processing functions for cropping faces."""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
@@ -10,9 +10,6 @@ from .config import Config
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# Lazy singleton
|
|
||||||
_yolo_model = None
|
|
||||||
|
|
||||||
|
|
||||||
def _save_jpeg(img: Image.Image, path: str) -> None:
|
def _save_jpeg(img: Image.Image, path: str) -> None:
|
||||||
if img.mode != "RGB":
|
if img.mode != "RGB":
|
||||||
@@ -20,17 +17,6 @@ def _save_jpeg(img: Image.Image, path: str) -> None:
|
|||||||
img.save(path, format="JPEG")
|
img.save(path, format="JPEG")
|
||||||
|
|
||||||
|
|
||||||
def get_yolo_model():
|
|
||||||
"""Singleton for YOLO model."""
|
|
||||||
global _yolo_model
|
|
||||||
if _yolo_model is None:
|
|
||||||
from ultralytics import YOLO
|
|
||||||
|
|
||||||
logger.info("Loading YOLOv9c model...")
|
|
||||||
_yolo_model = YOLO("yolov9c.pt")
|
|
||||||
return _yolo_model
|
|
||||||
|
|
||||||
|
|
||||||
def align_face(img: Image.Image, landmarks: list[list[float]] | np.ndarray) -> Image.Image | None:
|
def align_face(img: Image.Image, landmarks: list[list[float]] | np.ndarray) -> Image.Image | None:
|
||||||
"""Align face using 5-point landmarks to standard ArcFace input format (112x112).
|
"""Align face using 5-point landmarks to standard ArcFace input format (112x112).
|
||||||
|
|
||||||
@@ -170,49 +156,4 @@ def process_face_mode(
|
|||||||
return face_crop.size
|
return face_crop.size
|
||||||
|
|
||||||
|
|
||||||
def process_object_mode(
|
|
||||||
img: Image.Image,
|
|
||||||
config: dict,
|
|
||||||
output_dir: str,
|
|
||||||
count: int,
|
|
||||||
) -> bool:
|
|
||||||
"""Detect and crop objects using YOLO."""
|
|
||||||
try:
|
|
||||||
model = get_yolo_model()
|
|
||||||
target_class = config.get("object_class", "dog")
|
|
||||||
import torch
|
|
||||||
|
|
||||||
if os.getenv("FORCE_CPU", "").lower() in ("true", "1", "yes"):
|
|
||||||
device = "cpu"
|
|
||||||
elif hasattr(torch, "xpu") and torch.xpu.is_available():
|
|
||||||
device = "xpu"
|
|
||||||
else:
|
|
||||||
device = None # YOLO auto-selects (CUDA/ROCm/CPU)
|
|
||||||
|
|
||||||
results = model(img, verbose=False, device=device)
|
|
||||||
|
|
||||||
found = False
|
|
||||||
class_idx = 0 # Sequential counter per target class (Issue #10)
|
|
||||||
for box in (box for r in results for box in r.boxes):
|
|
||||||
cls_id = int(box.cls[0])
|
|
||||||
conf = float(box.conf[0])
|
|
||||||
if 0 <= cls_id < len(model.names) and model.names[cls_id] == target_class and conf > 0.5:
|
|
||||||
x1, y1, x2, y2 = box.xyxy[0].tolist()
|
|
||||||
_save_jpeg(
|
|
||||||
img.crop((x1, y1, x2, y2)),
|
|
||||||
os.path.join(output_dir, f"{count}_{class_idx}.jpg"),
|
|
||||||
)
|
|
||||||
class_idx += 1
|
|
||||||
found = True
|
|
||||||
|
|
||||||
return found
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"YOLO processing failed: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
def process_full_mode(img: Image.Image, output_dir: str, count: int) -> bool:
|
|
||||||
"""Save full image."""
|
|
||||||
_save_jpeg(img, os.path.join(output_dir, f"{count}.jpg"))
|
|
||||||
return True
|
|
||||||
|
|
||||||
|
|||||||
+15
-44
@@ -26,10 +26,8 @@ STRATEGY_PRESETS = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def _get_strategy_choice(has_embedding: bool, entity_type: str) -> tuple[int | str, str]:
|
def _get_strategy_choice(has_embedding: bool) -> tuple[int | str, str]:
|
||||||
"""Prompt user for training strategy and return (limit, selection_mode)."""
|
"""Prompt user for training strategy and return (limit, selection_mode)."""
|
||||||
model_name = "InsightFace" if entity_type == "face" else "SigLIP"
|
|
||||||
|
|
||||||
if has_embedding:
|
if has_embedding:
|
||||||
rprint(" [bold]1.[/bold] Adaptive Diversity [green][Recommended][/green]")
|
rprint(" [bold]1.[/bold] Adaptive Diversity [green][Recommended][/green]")
|
||||||
rprint(" [dim]• Dynamically selects images until redundancy starts[/dim]")
|
rprint(" [dim]• Dynamically selects images until redundancy starts[/dim]")
|
||||||
@@ -51,7 +49,7 @@ def _get_strategy_choice(has_embedding: bool, entity_type: str) -> tuple[int | s
|
|||||||
return 30, "smart"
|
return 30, "smart"
|
||||||
|
|
||||||
# Fallback when embedding model not available
|
# Fallback when embedding model not available
|
||||||
rprint(f" [yellow]Note: {model_name} not available. Using Time Spread.[/yellow]")
|
rprint(" [yellow]Note: InsightFace not available. Using Time Spread.[/yellow]")
|
||||||
rprint(" [bold]1.[/bold] Standard (30 images) [green][Recommended][/green]")
|
rprint(" [bold]1.[/bold] Standard (30 images) [green][Recommended][/green]")
|
||||||
rprint(" [bold]2.[/bold] Broad (100 images)")
|
rprint(" [bold]2.[/bold] Broad (100 images)")
|
||||||
rprint(" [bold]3.[/bold] Custom Count")
|
rprint(" [bold]3.[/bold] Custom Count")
|
||||||
@@ -86,15 +84,13 @@ def _resolve_strategy(strategy: str, has_embedding: bool) -> tuple[int | str, st
|
|||||||
|
|
||||||
|
|
||||||
def _perform_selection(
|
def _perform_selection(
|
||||||
assets: list, limit: int | str, name: str, selection_mode: str, entity_type: str, person_id: str | None = None
|
assets: list, limit: int | str, name: str, selection_mode: str, person_id: str | None = None
|
||||||
) -> list:
|
) -> list:
|
||||||
"""Run diversity selection with progress display."""
|
"""Run diversity selection with progress display."""
|
||||||
if selection_mode == "smart":
|
if selection_mode == "smart":
|
||||||
model_display = "InsightFace (face embeddings)" if entity_type == "face" else "SigLIP (visual embeddings)"
|
rprint("\n[cyan]Using InsightFace (face embeddings) for diversity analysis...[/cyan]")
|
||||||
rprint(f"\n[cyan]Using {model_display} for diversity analysis...[/cyan]")
|
|
||||||
|
|
||||||
# Pre-load model explicitly (separate from availability check)
|
load_embedding_model()
|
||||||
load_embedding_model(entity_type)
|
|
||||||
|
|
||||||
with Progress(
|
with Progress(
|
||||||
SpinnerColumn(),
|
SpinnerColumn(),
|
||||||
@@ -109,7 +105,6 @@ def _perform_selection(
|
|||||||
limit,
|
limit,
|
||||||
name,
|
name,
|
||||||
selection_mode=selection_mode,
|
selection_mode=selection_mode,
|
||||||
entity_type=entity_type,
|
|
||||||
person_id=person_id,
|
person_id=person_id,
|
||||||
progress_callback=lambda c, t: progress.update(task, completed=c, total=t),
|
progress_callback=lambda c, t: progress.update(task, completed=c, total=t),
|
||||||
)
|
)
|
||||||
@@ -120,9 +115,7 @@ def _perform_selection(
|
|||||||
|
|
||||||
rprint(f"\n[cyan]Using time-spread selection for {limit} images...[/cyan]")
|
rprint(f"\n[cyan]Using time-spread selection for {limit} images...[/cyan]")
|
||||||
with console.status(f"[bold]Selecting {limit} images evenly distributed over time...[/bold]"):
|
with console.status(f"[bold]Selecting {limit} images evenly distributed over time...[/bold]"):
|
||||||
selected = select_diverse_assets(
|
selected = select_diverse_assets(assets, limit, name, selection_mode="time", person_id=person_id)
|
||||||
assets, limit, name, selection_mode="time", entity_type=entity_type, person_id=person_id
|
|
||||||
)
|
|
||||||
rprint(f" [green]Selected {len(selected)} images using time spread.[/green]")
|
rprint(f" [green]Selected {len(selected)} images using time spread.[/green]")
|
||||||
return selected
|
return selected
|
||||||
|
|
||||||
@@ -132,22 +125,12 @@ def _configure_person(person: dict, people: list[dict]) -> dict | None:
|
|||||||
name = person["name"]
|
name = person["name"]
|
||||||
console.print(f"\nSelected: [bold green]{name}[/bold green]")
|
console.print(f"\nSelected: [bold green]{name}[/bold green]")
|
||||||
|
|
||||||
# Select training mode
|
config = {"name": name, "mode": "face", "quality_replacement": Config.QUALITY_REPLACEMENT}
|
||||||
rprint("\n[bold cyan]Training Mode:[/bold cyan]")
|
|
||||||
rprint(" [bold]1.[/bold] Face (Frigate Face Recognition)")
|
|
||||||
rprint(" [bold]2.[/bold] Object (Frigate Object Classification)")
|
|
||||||
|
|
||||||
mode_choice = Prompt.ask("Choice", choices=["1", "2"], default="1")
|
|
||||||
entity_type = "face" if mode_choice == "1" else "object"
|
|
||||||
|
|
||||||
config = {"name": name, "mode": entity_type, "quality_replacement": Config.QUALITY_REPLACEMENT}
|
|
||||||
if entity_type == "object":
|
|
||||||
config["object_class"] = Prompt.ask("Enter Object Class (e.g. dog, cat, car)", default="dog")
|
|
||||||
|
|
||||||
# Fetch and filter assets
|
# Fetch and filter assets
|
||||||
years = IntPrompt.ask("Filter images older than (years)", default=Config.YEARS_FILTER)
|
years = IntPrompt.ask("Filter images older than (years)", default=Config.YEARS_FILTER)
|
||||||
|
|
||||||
console.print(f"Scanning for {name} ({entity_type})...")
|
console.print(f"Scanning for {name}...")
|
||||||
with console.status("[bold green]Fetching assets...[/bold green]"):
|
with console.status("[bold green]Fetching assets...[/bold green]"):
|
||||||
all_assets = fetch_all_assets(person)
|
all_assets = fetch_all_assets(person)
|
||||||
recent_assets = filter_recent_assets(all_assets, years=years)
|
recent_assets = filter_recent_assets(all_assets, years=years)
|
||||||
@@ -171,17 +154,15 @@ def _configure_person(person: dict, people: list[dict]) -> dict | None:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
# Strategy selection
|
# Strategy selection
|
||||||
has_embedding = is_embedding_available(entity_type)
|
has_embedding = is_embedding_available()
|
||||||
rprint(f"\n[bold cyan]Select Training Strategy for {name}:[/bold cyan]")
|
rprint(f"\n[bold cyan]Select Training Strategy for {name}:[/bold cyan]")
|
||||||
|
|
||||||
limit, selection_mode = _get_strategy_choice(has_embedding, entity_type)
|
limit, selection_mode = _get_strategy_choice(has_embedding)
|
||||||
if selection_mode == "skip":
|
if selection_mode == "skip":
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# Perform selection
|
# Perform selection
|
||||||
selected_assets = _perform_selection(
|
selected_assets = _perform_selection(recent_assets, limit, name, selection_mode, person_id=person["id"])
|
||||||
recent_assets, limit, name, selection_mode, entity_type, person_id=person["id"]
|
|
||||||
)
|
|
||||||
|
|
||||||
rprint(f" [green]Queued {len(selected_assets)} images for {name}.[/green]")
|
rprint(f" [green]Queued {len(selected_assets)} images for {name}.[/green]")
|
||||||
return {"person": person, "assets": selected_assets, "limit": len(selected_assets), "config": config}
|
return {"person": person, "assets": selected_assets, "limit": len(selected_assets), "config": config}
|
||||||
@@ -231,7 +212,6 @@ def auto_configure(people: list[dict]) -> list[dict]:
|
|||||||
rprint("[red]No people found with names in Immich.[/red]")
|
rprint("[red]No people found with names in Immich.[/red]")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
mode = os.environ.get("TRAINING_MODE", "face")
|
|
||||||
strategy = os.environ.get("STRATEGY", "auto")
|
strategy = os.environ.get("STRATEGY", "auto")
|
||||||
skip = os.environ.get("SKIP_PEOPLE", "").split(",") if os.environ.get("SKIP_PEOPLE") else []
|
skip = os.environ.get("SKIP_PEOPLE", "").split(",") if os.environ.get("SKIP_PEOPLE") else []
|
||||||
only = os.environ.get("ONLY_PEOPLE", "").split(",") if os.environ.get("ONLY_PEOPLE") else []
|
only = os.environ.get("ONLY_PEOPLE", "").split(",") if os.environ.get("ONLY_PEOPLE") else []
|
||||||
@@ -260,11 +240,7 @@ def auto_configure(people: list[dict]) -> list[dict]:
|
|||||||
jobs = []
|
jobs = []
|
||||||
for person in valid_people:
|
for person in valid_people:
|
||||||
name = person["name"]
|
name = person["name"]
|
||||||
entity_type = mode
|
config = {"name": name, "mode": "face"}
|
||||||
|
|
||||||
config = {"name": name, "mode": entity_type}
|
|
||||||
if entity_type == "object":
|
|
||||||
config["object_class"] = os.environ.get("OBJECT_CLASS", "dog")
|
|
||||||
|
|
||||||
all_assets = fetch_all_assets(person)
|
all_assets = fetch_all_assets(person)
|
||||||
recent_assets = filter_recent_assets(all_assets, years=Config.YEARS_FILTER)
|
recent_assets = filter_recent_assets(all_assets, years=Config.YEARS_FILTER)
|
||||||
@@ -307,7 +283,7 @@ def auto_configure(people: list[dict]) -> list[dict]:
|
|||||||
|
|
||||||
config["quality_replacement"] = quality_replacement_only or Config.QUALITY_REPLACEMENT
|
config["quality_replacement"] = quality_replacement_only or Config.QUALITY_REPLACEMENT
|
||||||
|
|
||||||
has_embedding = is_embedding_available(entity_type)
|
has_embedding = is_embedding_available()
|
||||||
limit, selection_mode = _resolve_strategy(strategy, has_embedding)
|
limit, selection_mode = _resolve_strategy(strategy, has_embedding)
|
||||||
|
|
||||||
# Cap selection to remaining capacity (no cap when replacement-only — executor
|
# Cap selection to remaining capacity (no cap when replacement-only — executor
|
||||||
@@ -323,9 +299,7 @@ def auto_configure(people: list[dict]) -> list[dict]:
|
|||||||
if selection_mode == "skip":
|
if selection_mode == "skip":
|
||||||
continue
|
continue
|
||||||
|
|
||||||
selected_assets = _perform_selection(
|
selected_assets = _perform_selection(recent_assets, limit, name, selection_mode, person_id=person["id"])
|
||||||
recent_assets, limit, name, selection_mode, entity_type, person_id=person["id"]
|
|
||||||
)
|
|
||||||
if auto_cap is not None:
|
if auto_cap is not None:
|
||||||
selected_assets = selected_assets[:auto_cap]
|
selected_assets = selected_assets[:auto_cap]
|
||||||
|
|
||||||
@@ -340,20 +314,17 @@ def _show_preview(jobs: list[dict]) -> None:
|
|||||||
"""Show a summary table of all queued jobs before execution."""
|
"""Show a summary table of all queued jobs before execution."""
|
||||||
table = Table(title="📋 Training Job Preview", show_header=True, header_style="bold cyan")
|
table = Table(title="📋 Training Job Preview", show_header=True, header_style="bold cyan")
|
||||||
table.add_column("Person", style="bold")
|
table.add_column("Person", style="bold")
|
||||||
table.add_column("Mode", style="dim")
|
|
||||||
table.add_column("Images", justify="right")
|
table.add_column("Images", justify="right")
|
||||||
table.add_column("Date Range", style="dim")
|
table.add_column("Date Range", style="dim")
|
||||||
|
|
||||||
for job in jobs:
|
for job in jobs:
|
||||||
name = job["person"]["name"]
|
name = job["person"]["name"]
|
||||||
mode = job["config"].get("mode", "face")
|
|
||||||
count = str(job["limit"])
|
count = str(job["limit"])
|
||||||
|
|
||||||
# Date range
|
|
||||||
dates = sorted(a.get("fileCreatedAt", "")[:10] for a in job["assets"] if a.get("fileCreatedAt"))
|
dates = sorted(a.get("fileCreatedAt", "")[:10] for a in job["assets"] if a.get("fileCreatedAt"))
|
||||||
date_range = f"{dates[0]} → {dates[-1]}" if len(dates) >= 2 else (dates[0] if dates else "—")
|
date_range = f"{dates[0]} → {dates[-1]}" if len(dates) >= 2 else (dates[0] if dates else "—")
|
||||||
|
|
||||||
table.add_row(name, mode, count, date_range)
|
table.add_row(name, count, date_range)
|
||||||
|
|
||||||
console.print()
|
console.print()
|
||||||
console.print(table)
|
console.print(table)
|
||||||
|
|||||||
Reference in New Issue
Block a user