Files
frigate/frigate/embeddings/__init__.py
T

331 lines
11 KiB
Python
Raw Normal View History

"""SQLite-vec embeddings database."""
2024-06-21 17:30:19 -04:00
2024-10-22 16:05:48 -06:00
import base64
2024-06-23 09:13:02 -04:00
import json
2024-06-21 17:30:19 -04:00
import logging
import os
2026-04-29 17:20:19 -05:00
import sys
2024-06-21 17:30:19 -04:00
import threading
2025-07-29 13:38:13 -05:00
from json.decoder import JSONDecodeError
2025-06-24 11:41:11 -06:00
from multiprocessing.synchronize import Event as MpEvent
2025-06-12 12:12:34 -06:00
from typing import Any, Union
2024-06-21 17:30:19 -04:00
2025-05-17 17:11:19 -05:00
import regex
from pathvalidate import ValidationError, sanitize_filename
2024-06-21 17:30:19 -04:00
2024-10-10 09:42:24 -06:00
from frigate.comms.embeddings_updater import EmbeddingsRequestEnum, EmbeddingsRequestor
2024-06-21 17:30:19 -04:00
from frigate.config import FrigateConfig
from frigate.const import CONFIG_DIR, FACE_DIR, PROCESS_PRIORITY_HIGH
2025-01-10 12:44:30 -07:00
from frigate.data_processing.types import DataProcessorMetrics
from frigate.db.sqlitevecq import SqliteVecQueueDatabase
2025-06-12 12:12:34 -06:00
from frigate.models import Event
2024-10-10 09:42:24 -06:00
from frigate.util.builtin import serialize
2025-06-09 08:25:33 -06:00
from frigate.util.classification import kickoff_model_training
2025-06-13 11:09:51 -06:00
from frigate.util.process import FrigateProcess
2024-06-21 17:30:19 -04:00
2024-06-23 09:13:02 -04:00
from .maintainer import EmbeddingMaintainer
from .util import ZScoreNormalization
2024-06-21 17:30:19 -04:00
logger = logging.getLogger(__name__)
2025-06-12 12:12:34 -06:00
class EmbeddingProcess(FrigateProcess):
def __init__(
2025-06-24 11:41:11 -06:00
self,
config: FrigateConfig,
metrics: DataProcessorMetrics | None,
stop_event: MpEvent,
2025-06-12 12:12:34 -06:00
) -> None:
super().__init__(
stop_event,
PROCESS_PRIORITY_HIGH,
name="frigate.embeddings_manager",
daemon=True,
)
2025-06-12 12:12:34 -06:00
self.config = config
self.metrics = metrics
2024-06-21 17:30:19 -04:00
2025-06-12 12:12:34 -06:00
def run(self) -> None:
self.pre_run_setup(self.config.logger)
2025-06-12 12:12:34 -06:00
maintainer = EmbeddingMaintainer(
self.config,
self.metrics,
self.stop_event,
)
maintainer.start()
2026-04-29 17:20:19 -05:00
maintainer.join()
# If the maintainer thread exited but no shutdown was requested, it
# crashed. Surface as a non-zero exit so the watchdog restarts us
# instead of treating the silent thread death as a clean shutdown.
if not self.stop_event.is_set():
logger.error("Embeddings maintainer thread exited unexpectedly")
sys.exit(1)
2024-06-23 09:13:02 -04:00
class EmbeddingsContext:
2024-10-10 09:42:24 -06:00
def __init__(self, db: SqliteVecQueueDatabase):
self.db = db
2024-06-23 09:13:02 -04:00
self.thumb_stats = ZScoreNormalization()
2024-10-09 16:31:54 -05:00
self.desc_stats = ZScoreNormalization()
2024-10-10 09:42:24 -06:00
self.requestor = EmbeddingsRequestor()
2024-06-23 09:13:02 -04:00
# load stats from disk
2025-07-29 13:38:13 -05:00
stats_file = os.path.join(CONFIG_DIR, ".search_stats.json")
2024-06-23 09:13:02 -04:00
try:
2025-07-29 13:38:13 -05:00
with open(stats_file, "r") as f:
2024-06-23 09:13:02 -04:00
data = json.loads(f.read())
self.thumb_stats.from_dict(data["thumb_stats"])
self.desc_stats.from_dict(data["desc_stats"])
except FileNotFoundError:
pass
2025-07-29 13:38:13 -05:00
except JSONDecodeError:
logger.warning("Failed to decode semantic search stats, clearing file")
try:
with open(stats_file, "w") as f:
f.write("")
except OSError as e:
logger.error(f"Failed to clear corrupted stats file: {e}")
2024-06-23 09:13:02 -04:00
2024-10-10 09:42:24 -06:00
def stop(self):
2024-06-23 09:13:02 -04:00
"""Write the stats to disk as JSON on exit."""
contents = {
"thumb_stats": self.thumb_stats.to_dict(),
"desc_stats": self.desc_stats.to_dict(),
}
with open(os.path.join(CONFIG_DIR, ".search_stats.json"), "w") as f:
json.dump(contents, f)
2024-10-10 09:42:24 -06:00
self.requestor.stop()
def search_thumbnail(
self, query: Union[Event, str], event_ids: list[str] = None
) -> list[tuple[str, float]]:
if query.__class__ == Event:
cursor = self.db.execute_sql(
"""
SELECT thumbnail_embedding FROM vec_thumbnails WHERE id = ?
""",
[query.id],
)
row = cursor.fetchone() if cursor else None
if row:
query_embedding = row[0]
else:
# If no embedding found, generate it and return it
2024-10-10 15:37:43 -06:00
data = self.requestor.send_data(
EmbeddingsRequestEnum.embed_thumbnail.value,
{"id": str(query.id), "thumbnail": str(query.thumbnail)},
2024-10-10 09:42:24 -06:00
)
2024-10-10 15:37:43 -06:00
if not data:
return []
query_embedding = serialize(data)
2024-10-10 09:42:24 -06:00
else:
2024-10-10 15:37:43 -06:00
data = self.requestor.send_data(
EmbeddingsRequestEnum.generate_search.value, query
2024-10-10 09:42:24 -06:00
)
2024-10-10 15:37:43 -06:00
if not data:
return []
query_embedding = serialize(data)
2024-10-10 09:42:24 -06:00
sql_query = """
SELECT
id,
distance
FROM vec_thumbnails
WHERE thumbnail_embedding MATCH ?
AND k = 100
"""
# Add the IN clause if event_ids is provided and not empty
# this is the only filter supported by sqlite-vec as of 0.1.3
# but it seems to be broken in this version
if event_ids:
sql_query += " AND id IN ({})".format(",".join("?" * len(event_ids)))
# order by distance DESC is not implemented in this version of sqlite-vec
# when it's implemented, we can use cosine similarity
sql_query += " ORDER BY distance"
parameters = [query_embedding] + event_ids if event_ids else [query_embedding]
results = self.db.execute_sql(sql_query, parameters).fetchall()
return results
def search_description(
self, query_text: str, event_ids: list[str] = None
) -> list[tuple[str, float]]:
2024-10-10 15:37:43 -06:00
data = self.requestor.send_data(
EmbeddingsRequestEnum.generate_search.value, query_text
2024-10-10 09:42:24 -06:00
)
2024-10-10 15:37:43 -06:00
if not data:
return []
query_embedding = serialize(data)
2024-10-10 09:42:24 -06:00
# Prepare the base SQL query
sql_query = """
SELECT
id,
distance
FROM vec_descriptions
WHERE description_embedding MATCH ?
AND k = 100
"""
# Add the IN clause if event_ids is provided and not empty
# this is the only filter supported by sqlite-vec as of 0.1.3
# but it seems to be broken in this version
if event_ids:
sql_query += " AND id IN ({})".format(",".join("?" * len(event_ids)))
# order by distance DESC is not implemented in this version of sqlite-vec
# when it's implemented, we can use cosine similarity
sql_query += " ORDER BY distance"
parameters = [query_embedding] + event_ids if event_ids else [query_embedding]
results = self.db.execute_sql(sql_query, parameters).fetchall()
return results
2025-05-13 16:27:20 +02:00
def register_face(self, face_name: str, image_data: bytes) -> dict[str, Any]:
2025-01-10 08:39:24 -07:00
return self.requestor.send_data(
2024-10-22 16:05:48 -06:00
EmbeddingsRequestEnum.register_face.value,
{
"face_name": face_name,
"image": base64.b64encode(image_data).decode("ASCII"),
},
)
2025-05-13 16:27:20 +02:00
def recognize_face(self, image_data: bytes) -> dict[str, Any]:
2025-03-19 09:02:25 -06:00
return self.requestor.send_data(
EmbeddingsRequestEnum.recognize_face.value,
{
"image": base64.b64encode(image_data).decode("ASCII"),
},
)
2024-10-22 16:05:48 -06:00
def get_face_ids(self, name: str) -> list[str]:
sql_query = """
2024-10-22 16:05:48 -06:00
SELECT
id
FROM vec_descriptions
WHERE id LIKE ?
2024-10-22 16:05:48 -06:00
"""
return self.db.execute_sql(sql_query, (f"%{name}%",)).fetchall()
2024-10-22 16:05:48 -06:00
2025-05-13 16:27:20 +02:00
def reprocess_face(self, face_file: str) -> dict[str, Any]:
2025-01-29 07:41:35 -07:00
return self.requestor.send_data(
EmbeddingsRequestEnum.reprocess_face.value, {"image_file": face_file}
)
2025-01-18 10:52:01 -07:00
def clear_face_classifier(self) -> None:
self.requestor.send_data(
EmbeddingsRequestEnum.clear_face_classifier.value, None
)
2024-11-26 13:41:49 -07:00
def delete_face_ids(self, face: str, ids: list[str]) -> None:
folder = os.path.join(FACE_DIR, face)
for id in ids:
file_path = os.path.join(folder, id)
if os.path.isfile(file_path):
os.unlink(file_path)
2024-10-22 16:05:48 -06:00
2025-05-09 08:36:44 -05:00
if face != "train" and len(os.listdir(folder)) == 0:
2025-03-16 05:01:15 -06:00
os.rmdir(folder)
self.requestor.send_data(
EmbeddingsRequestEnum.clear_face_classifier.value, None
)
def rename_face(self, old_name: str, new_name: str) -> None:
2025-05-17 17:11:19 -05:00
valid_name_pattern = r"^[\p{L}\p{N}\s'_-]{1,50}$"
try:
sanitized_old_name = sanitize_filename(old_name, replacement_text="_")
sanitized_new_name = sanitize_filename(new_name, replacement_text="_")
except ValidationError as e:
raise ValueError(f"Invalid face name: {str(e)}")
2025-05-17 17:11:19 -05:00
if not regex.match(valid_name_pattern, old_name):
raise ValueError(f"Invalid old face name: {old_name}")
2025-05-17 17:11:19 -05:00
if not regex.match(valid_name_pattern, new_name):
raise ValueError(f"Invalid new face name: {new_name}")
if sanitized_old_name != old_name:
raise ValueError(f"Old face name contains invalid characters: {old_name}")
if sanitized_new_name != new_name:
raise ValueError(f"New face name contains invalid characters: {new_name}")
old_path = os.path.normpath(os.path.join(FACE_DIR, old_name))
new_path = os.path.normpath(os.path.join(FACE_DIR, new_name))
# Prevent path traversal
if not old_path.startswith(
os.path.normpath(FACE_DIR)
) or not new_path.startswith(os.path.normpath(FACE_DIR)):
raise ValueError("Invalid path detected")
if not os.path.exists(old_path):
raise ValueError(f"Face {old_name} not found.")
os.rename(old_path, new_path)
self.requestor.send_data(
EmbeddingsRequestEnum.clear_face_classifier.value, None
)
2024-10-10 09:42:24 -06:00
def update_description(self, event_id: str, description: str) -> None:
self.requestor.send_data(
EmbeddingsRequestEnum.embed_description.value,
{"id": event_id, "description": description},
)
2025-05-13 16:27:20 +02:00
def reprocess_plate(self, event: dict[str, Any]) -> dict[str, Any]:
return self.requestor.send_data(
EmbeddingsRequestEnum.reprocess_plate.value, {"event": event}
)
2025-03-27 12:29:34 -05:00
2025-05-13 16:27:20 +02:00
def reindex_embeddings(self) -> dict[str, Any]:
2025-03-27 12:29:34 -05:00
return self.requestor.send_data(EmbeddingsRequestEnum.reindex.value, {})
2025-05-27 10:26:00 -05:00
2025-06-05 09:13:12 -06:00
def start_classification_training(self, model_name: str) -> dict[str, Any]:
2025-06-09 08:25:33 -06:00
threading.Thread(
target=kickoff_model_training,
args=(self.requestor, model_name),
daemon=True,
).start()
return {"success": True, "message": f"Began training {model_name} model."}
2025-06-05 09:13:12 -06:00
2025-05-27 10:26:00 -05:00
def transcribe_audio(self, event: dict[str, any]) -> dict[str, any]:
return self.requestor.send_data(
EmbeddingsRequestEnum.transcribe_audio.value, {"event": event}
)
2025-07-07 09:03:57 -05:00
def generate_description_embedding(self, text: str) -> None:
return self.requestor.send_data(
EmbeddingsRequestEnum.embed_description.value,
{"id": None, "description": text, "upsert": False},
)
def generate_image_embedding(self, event_id: str, thumbnail: bytes) -> None:
return self.requestor.send_data(
EmbeddingsRequestEnum.embed_thumbnail.value,
{"id": str(event_id), "thumbnail": str(thumbnail), "upsert": False},
)
2025-08-12 16:27:35 -06:00
def generate_review_summary(self, start_ts: float, end_ts: float) -> str | None:
return self.requestor.send_data(
EmbeddingsRequestEnum.summarize_review.value,
{"start_ts": start_ts, "end_ts": end_ts},
)