-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdatabase.py
More file actions
125 lines (104 loc) · 4.59 KB
/
Copy pathdatabase.py
File metadata and controls
125 lines (104 loc) · 4.59 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
"""Vector storage for speaker embeddings (NumPy + JSON labels)."""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any, Optional, Union
import numpy as np
DEFAULT_DB_DIR = Path(__file__).resolve().parent / "speaker_db"
EMBEDDINGS_FILE = "embeddings.npy"
LABELS_FILE = "labels.json"
class SpeakerDatabase:
"""Stores L2-normalized embeddings and an ordered list of speaker names."""
def __init__(self, db_dir: Optional[Union[Path, str]] = None) -> None:
self.db_dir = Path(db_dir) if db_dir else DEFAULT_DB_DIR
self.embeddings: Optional[np.ndarray] = None
self.labels: list[str] = []
@property
def embeddings_path(self) -> Path:
return self.db_dir / EMBEDDINGS_FILE
@property
def labels_path(self) -> Path:
return self.db_dir / LABELS_FILE
def load(self) -> None:
"""Load database from disk; starts empty if files are missing."""
self.db_dir.mkdir(parents=True, exist_ok=True)
if not self.embeddings_path.exists() or not self.labels_path.exists():
self.embeddings = None
self.labels = []
return
self.embeddings = np.load(self.embeddings_path)
with self.labels_path.open(encoding="utf-8") as f:
payload: dict[str, Any] = json.load(f)
names = payload.get("labels")
if not isinstance(names, list) or not all(isinstance(x, str) for x in names):
raise ValueError("labels.json must contain a JSON object with key 'labels'.")
self.labels = names
if self.embeddings.shape[0] != len(self.labels):
raise ValueError(
"embeddings.npy row count must match len(labels). "
f"Got {self.embeddings.shape[0]} vs {len(self.labels)}."
)
def save(self) -> None:
"""Persist embeddings and labels to ``db_dir``."""
self.db_dir.mkdir(parents=True, exist_ok=True)
if self.embeddings is None or self.embeddings.size == 0:
if self.embeddings_path.exists():
self.embeddings_path.unlink()
if self.labels_path.exists():
self.labels_path.unlink()
return
np.save(self.embeddings_path, self.embeddings.astype(np.float32))
with self.labels_path.open("w", encoding="utf-8") as f:
json.dump({"labels": self.labels}, f, indent=2)
def is_empty(self) -> bool:
return self.embeddings is None or self.embeddings.shape[0] == 0
def add_speaker(self, name: str, embedding: np.ndarray) -> None:
"""Append one normalized embedding for ``name``."""
row = np.asarray(embedding, dtype=np.float32).reshape(1, -1)
if self.embeddings is None:
self.embeddings = row
else:
if row.shape[1] != self.embeddings.shape[1]:
raise ValueError(
f"Embedding dim {row.shape[1]} != db dim {self.embeddings.shape[1]}."
)
self.embeddings = np.vstack([self.embeddings, row])
self.labels.append(name)
def replace_speaker(self, name: str, embedding: np.ndarray) -> None:
"""Remove existing rows for ``name`` and add one new embedding."""
if not self.labels:
self.add_speaker(name, embedding)
return
emb = np.asarray(embedding, dtype=np.float32).reshape(1, -1)
indices_to_keep: list[int] = [
i for i, label in enumerate(self.labels) if label != name
]
if not indices_to_keep:
self.embeddings = emb
self.labels = [name]
return
self.embeddings = np.vstack([self.embeddings[indices_to_keep], emb])
self.labels = [self.labels[i] for i in indices_to_keep] + [name]
def delete_speaker(self, name: str) -> bool:
"""Delete all rows for ``name`` and return whether anything changed."""
if not self.labels:
return False
indices_to_keep: list[int] = [
i for i, label in enumerate(self.labels) if label != name
]
if len(indices_to_keep) == len(self.labels):
return False
if not indices_to_keep:
self.embeddings = None
self.labels = []
return True
assert self.embeddings is not None
self.embeddings = self.embeddings[indices_to_keep]
self.labels = [self.labels[i] for i in indices_to_keep]
return True
def load_db(db_dir: Optional[Union[Path, str]] = None) -> SpeakerDatabase:
db = SpeakerDatabase(db_dir=db_dir)
db.load()
return db
def save_db(db: SpeakerDatabase) -> None:
db.save()