diff --git a/.github/workflows/python-package.yml b/.github/workflows/python-package.yml index 55514aa..5da41b2 100644 --- a/.github/workflows/python-package.yml +++ b/.github/workflows/python-package.yml @@ -43,4 +43,4 @@ jobs: uses: actions/upload-artifact@v4 with: name: coverage-${{ matrix.python-version }} - path: .coverage.xml + path: coverage.xml diff --git a/src/vocalid/audio_utils.py b/src/vocalid/audio_utils.py index 6227b1b..f0813ee 100644 --- a/src/vocalid/audio_utils.py +++ b/src/vocalid/audio_utils.py @@ -1,62 +1,156 @@ + + +import logging + import torch import torchaudio import sounddevice as sd -import numpy as np +import soundfile as sf + from .config import SAMPLE_RATE +logger = logging.getLogger(__name__) + def load_audio(path, target_sr=SAMPLE_RATE): - waveform, sr = torchaudio.load(path) + """ + Load a WAV file and return a mono waveform tensor. - # Convert to float32 - waveform = waveform.float() + Parameters + ---------- + path : str + Path to the WAV file. + target_sr : int + Desired sample rate. - # Resample if needed - if sr != target_sr: - waveform = torchaudio.functional.resample(waveform, sr, target_sr) + Returns + ------- + tuple(torch.Tensor, int) + Waveform with shape (1, T) and sample rate. + """ + + samples, sr = sf.read( + path, + dtype="float32", + always_2d=True + ) - # Mono - if waveform.ndim == 2 and waveform.shape[0] > 1: + waveform = torch.from_numpy(samples.T).float() + + # Convert stereo to mono + if waveform.shape[0] > 1: waveform = waveform.mean(dim=0, keepdim=True) - # Ensure shape is (1, T) + # Resample if necessary + if sr != target_sr: + waveform = torchaudio.functional.resample( + waveform, + sr, + target_sr + ) + sr = target_sr + + # Ensure shape (1, T) if waveform.ndim == 1: waveform = waveform.unsqueeze(0) - # Pad short audio - min_len = target_sr # 1 second minimum - if waveform.shape[1] < min_len: - pad_len = min_len - waveform.shape[1] - waveform = torch.nn.functional.pad(waveform, (0, pad_len)) + # Pad to at least one second + if waveform.shape[1] < target_sr: + pad_len = target_sr - waveform.shape[1] + waveform = torch.nn.functional.pad( + waveform, + (0, pad_len) + ) - return waveform + return waveform, sr + + +def record_audio( + duration=4, + sample_rate=SAMPLE_RATE +): + """ + Record audio from the microphone. -def record_audio(seconds=4, fs=SAMPLE_RATE): - print(f"Recording {seconds} seconds...") + Parameters + ---------- + duration : float + Recording duration in seconds. + sample_rate : int + Recording sample rate. + + Returns + ------- + torch.Tensor + Recorded waveform with shape (1, T). + """ + + logger.info("Recording %.1f seconds...", duration) - # Check if sounddevice sees any input devices try: devices = sd.query_devices() - has_input = any(d["max_input_channels"] > 0 for d in devices) - if not has_input: - raise RuntimeError("No input microphone found. Live recording cannot run here.") - except Exception: - raise RuntimeError("No audio interface available. Live recording unsupported.") - # Try to record + if not any(device["max_input_channels"] > 0 for device in devices): + raise RuntimeError("No microphone found.") + + except Exception as e: + raise RuntimeError( + "Unable to access an input microphone." + ) from e + try: - audio = sd.rec(int(seconds * fs), samplerate=fs, channels=1, dtype='float32') + audio = sd.rec( + int(duration * sample_rate), + samplerate=sample_rate, + channels=1, + dtype="float32" + ) + sd.wait() + except Exception as e: - raise RuntimeError("Live recording failed. Likely running in Colab.") from e + raise RuntimeError( + "Audio recording failed." + ) from e + + logger.info("Recording complete") + + waveform = torch.from_numpy(audio.T).float() + + # Ensure minimum duration of one second + if waveform.shape[1] < sample_rate: + pad_len = sample_rate - waveform.shape[1] + waveform = torch.nn.functional.pad( + waveform, + (0, pad_len) + ) + + return waveform + - print("Recording complete") +def save_audio( + waveform, + path, + sample_rate=SAMPLE_RATE +): + """ + Save a waveform tensor as a WAV file. - audio_tensor = torch.tensor(audio.T, dtype=torch.float32) + Parameters + ---------- + waveform : torch.Tensor + Audio waveform of shape (1, T) or (T,). + path : str + Output WAV file path. + sample_rate : int + Sample rate. + """ - # Ensure minimum length - if audio_tensor.shape[1] < fs: - pad_len = fs - audio_tensor.shape[1] - audio_tensor = torch.nn.functional.pad(audio_tensor, (0, pad_len)) + if waveform.ndim == 2: + waveform = waveform.squeeze(0) - return audio_tensor + sf.write( + path, + waveform.cpu().numpy(), + sample_rate + ) \ No newline at end of file diff --git a/src/vocalid/auth_adapter.py b/src/vocalid/auth_adapter.py new file mode 100644 index 0000000..58e0b3f --- /dev/null +++ b/src/vocalid/auth_adapter.py @@ -0,0 +1,48 @@ +""" +auth_adapter.py +Adapts VocalID's verification result into an authentication decision. +Deliberately has no ML in it - it just calls VoiceVerifier and +translates (verified, score) into an AuthenticationResult that +application code (door lock, login screen, attendance app, etc.) +can branch on. +""" + +from dataclasses import dataclass +from .verifier import VoiceVerifier +from . import config + + +@dataclass +class AuthenticationResult: + success: bool + confidence: float + + def __str__(self): + status = "ACCESS GRANTED" if self.success else "ACCESS DENIED" + return f"{status} (confidence: {self.confidence:.2f})" + + +class VoiceAuthenticator: + def __init__(self, model_path: str): + self.verifier = VoiceVerifier(model_path) + + def authenticate_file(self, audio_path: str) -> AuthenticationResult: + verified, score = self.verifier.verify_file(audio_path) + return AuthenticationResult(success=verified, confidence=score) + + def authenticate_live(self, seconds: float = None) -> AuthenticationResult: + from .recorder import record_one + + seconds = seconds or 4.0 + + audio = record_one( + seconds, + config.SAMPLE_RATE + ) + + verified, score = self.verifier.verify_array(audio) + + return AuthenticationResult( + success=verified, + confidence=score + ) \ No newline at end of file diff --git a/src/vocalid/cli.py b/src/vocalid/cli.py index a84fa3a..29b00d7 100644 --- a/src/vocalid/cli.py +++ b/src/vocalid/cli.py @@ -1,9 +1,21 @@ + + import argparse +import logging import os from glob import glob + +from .audio_utils import record_audio from .trainer import VoiceTrainer from .verifier import VoiceVerifier -from .audio_utils import record_audio + +logging.basicConfig( + level=logging.INFO, + format="%(message)s", +) + +logger = logging.getLogger(__name__) + def main(): parser = argparse.ArgumentParser(description="Voice Verifier CLI") @@ -11,34 +23,36 @@ def main(): # training command train = sub.add_parser("train", help="Train a voice authentication model") - train.add_argument("--positive", required=True, help = "Folder with your voice samples") - train.add_argument("--negative", required=True, help = "Folder with other voices") - train.add_argument("--output", default="voice_auth.pkl", help="path to save model") + train.add_argument("--positive", required=True, help="Folder with your voice samples") + train.add_argument("--negative", required=True, help="Folder with other voices") + train.add_argument("--output", default="voice_auth.pkl", help="Path to save model") # evaluate the model - evaluate = sub.add_parser("evaluate", help = "Evaluate the trained model") - evaluate.add_argument("--model", required=True, help= "Path to trained model") - evaluate.add_argument("--positive", required=True, help = "Folder with your voice samples") - evaluate.add_argument("--negative", required=True, help= "Folder with negative/random samples") + evaluate = sub.add_parser("evaluate", help="Evaluate the trained model") + evaluate.add_argument("--model", required=True, help="Path to trained model") + evaluate.add_argument("--positive", required=True, help="Folder with your voice samples") + evaluate.add_argument("--negative", required=True, help="Folder with negative/random samples") # verify file command - verify = sub.add_parser("verify", help = "Verify a voice file") - verify.add_argument("file", help="path to .wav voice file to verify") - verify.add_argument("--model", default="voice_auth.pkl", help="trained model path") + verify = sub.add_parser("verify", help="Verify a voice file") + verify.add_argument("file", help="Path to .wav voice file to verify") + verify.add_argument("--model", default="voice_auth.pkl", help="Trained model path") # live verification command - live = sub.add_parser("live", help = "Live microphone verification") - live.add_argument("--model", default="voice_auth.pkl", help= "trained model path") - live.add_argument("--seconds", type=int, default=4, help= "Recording duration in seconds") + live = sub.add_parser("live", help="Live microphone verification") + live.add_argument("--model", default="voice_auth.pkl", help="Trained model path") + live.add_argument("--seconds", type=int, default=4, help="Recording duration in seconds") args = parser.parse_args() if args.commands == "train": pos_files = glob(os.path.join(args.positive, "*.wav")) neg_files = glob(os.path.join(args.negative, "*.wav")) + trainer = VoiceTrainer() trainer.train(pos_files, neg_files, args.output) - print(f"Model saved to {args.output}") + + logger.info("Model saved to %s", args.output) elif args.commands == "evaluate": pos_files = glob(os.path.join(args.positive, "*.wav")) @@ -46,29 +60,37 @@ def main(): trainer = VoiceTrainer() trainer.load(args.model) - # X, y = trainer.prepare_features(pos_files, neg_files) + results = trainer.evaluate(pos_files, neg_files) - print("\n===== Evaluation Results =====") - print("\nAccuracy", round(results["accuracy"], 4)) - print("\nClassification Report:\n") - print(results["report"]) + logger.info("") + logger.info("===== Evaluation Results =====") + logger.info("") + logger.info("Accuracy: %.4f", results["accuracy"]) + logger.info("") + logger.info("Classification Report:") + logger.info("%s", results["report"]) elif args.commands == "verify": verifier = VoiceVerifier(args.model) + ok, score = verifier.verify_file(args.file) - print(f"Verified: {ok}, Score: {score:.2f}") + + logger.info("Verified: %s, Score: %.2f", ok, score) elif args.commands == "live": verifier = VoiceVerifier(args.model) + try: audio_tensor = record_audio(args.seconds) + except RuntimeError as e: - print(str(e)) + logger.error("%s", e) return ok, score = verifier.verify_array(audio_tensor) - print(f"Verified: {ok}, Score: {score:.2f}") + + logger.info("Verified: %s, Score: %.2f", ok, score) else: parser.print_help() \ No newline at end of file diff --git a/src/vocalid/dataset_manager.py b/src/vocalid/dataset_manager.py new file mode 100644 index 0000000..89cc118 --- /dev/null +++ b/src/vocalid/dataset_manager.py @@ -0,0 +1,59 @@ +""" +dataset_manager.py +Keeps the on-disk dataset in the layout trainer.py expects: + + dataset/ + my_voice/ (positive samples) + other_voices/ (negative samples) + +Just handles copying accepted files in with sequential, collision-free +names - no metadata, no database. +""" +# src/vocalid/typing_compat.py + +from typing import ( + List, + Dict, + Tuple, + Set, + Optional, + Union, + Any, + Callable, +) +import os +import shutil + +POSITIVE_DIR = "my_voice" +NEGATIVE_DIR = "other_voices" + + +class DatasetManager: + def __init__(self, root: str = "dataset"): + self.root = root + self.positive_dir = os.path.join(root, POSITIVE_DIR) + self.negative_dir = os.path.join(root, NEGATIVE_DIR) + os.makedirs(self.positive_dir, exist_ok=True) + os.makedirs(self.negative_dir, exist_ok=True) + + def _target_dir(self, label: str) -> str: + if label not in ("positive", "negative"): + raise ValueError("label must be 'positive' or 'negative'") + return self.positive_dir if label == "positive" else self.negative_dir + + def list_samples(self, label: str) -> List[str]: + target_dir = self._target_dir(label) + return [ + os.path.join(target_dir, f) + for f in sorted(os.listdir(target_dir)) + if f.lower().endswith(".wav") + ] + + def add_sample(self, src_path: str, label: str) -> str: + """Copies src_path into the dataset with the next free index. Returns the new path.""" + target_dir = self._target_dir(label) + existing = self.list_samples(label) + next_index = len(existing) + 1 + dest_path = os.path.join(target_dir, f"sample{next_index:03d}.wav") + shutil.copy2(src_path, dest_path) + return dest_path diff --git a/src/vocalid/embeddings.py b/src/vocalid/embeddings.py index 55df8c5..0e33116 100644 --- a/src/vocalid/embeddings.py +++ b/src/vocalid/embeddings.py @@ -1,25 +1,37 @@ + + +import logging + +import numpy as np import torch + from .audio_utils import load_audio -import numpy as np + +logger = logging.getLogger(__name__) class EmbeddingExtractor: def __init__(self, model_path="speechbrain/spkrec-ecapa-voxceleb"): - # lazy torchaudio patch + # Lazy torchaudio compatibility patch try: import torchaudio + if not hasattr(torchaudio, "list_audio_backends"): torchaudio.list_audio_backends = lambda: ["sox_io"] + except Exception: + # Older/newer torchaudio versions may not expose this API. + # Safe to continue because it is only a compatibility patch. pass try: from speechbrain.inference import EncoderClassifier + except Exception as e: raise ImportError( - "SpeechBrain could not load. Torchaudio is incompatible.\n" + "SpeechBrain could not load.\n" f"Original error: {e}" - ) + ) from e self.model = EncoderClassifier.from_hparams( source=model_path, @@ -27,53 +39,106 @@ def __init__(self, model_path="speechbrain/spkrec-ecapa-voxceleb"): savedir="pretrained_models/ecapa", ) - def _prepare_waveform(self, wav): - # Convert numpy → tensor - if not isinstance(wav, torch.Tensor): - wav = torch.tensor(wav, dtype=torch.float32) + def _prepare_waveform(self, waveform): + """ + Convert waveform into the format expected by SpeechBrain. + + Expected shape: + (1, T) + """ - # Ensure shape (1, T) - if wav.ndim == 1: - wav = wav.unsqueeze(0) + if not isinstance(waveform, torch.Tensor): + waveform = torch.tensor( + waveform, + dtype=torch.float32 + ) + + # Convert (T,) -> (1, T) + if waveform.ndim == 1: + waveform = waveform.unsqueeze(0) - # If multichannel, average - if wav.shape[0] > 1: - wav = wav.mean(dim=0, keepdim=True) + # Stereo -> Mono + if waveform.shape[0] > 1: + waveform = waveform.mean( + dim=0, + keepdim=True + ) - # Pad waves shorter than 1 sec (ECAPA expects enough frames) - min_len = 16000 # 1 s at 16kHz - if wav.shape[1] < min_len: - pad_len = min_len - wav.shape[1] - wav = torch.nn.functional.pad(wav, (0, pad_len)) + # Pad to at least 1 second + min_len = 16000 - return wav + if waveform.shape[1] < min_len: + waveform = torch.nn.functional.pad( + waveform, + (0, min_len - waveform.shape[1]) + ) + + return waveform + + def _normalize(self, embedding): + embedding = np.asarray(embedding).squeeze() + + norm = np.linalg.norm(embedding) - def _normalize(self, emb): - emb = np.asarray(emb).squeeze() - norm = np.linalg.norm(emb) if norm == 0: - return emb - return emb / norm + return embedding + + return embedding / norm def embed_file(self, path): + """ + Load a WAV file and return a normalized speaker embedding. + """ + try: - waveform = load_audio(path) + waveform, _ = load_audio(path) waveform = self._prepare_waveform(waveform) - except Exception: - return None - with torch.no_grad(): - emb = self.model.encode_batch(waveform)[0].cpu().numpy() + with torch.no_grad(): + embedding = ( + self.model + .encode_batch(waveform) + .squeeze() + .cpu() + .numpy() + ) - return self._normalize(emb) + return self._normalize(embedding) + + except Exception: + logger.exception( + "Embedding extraction failed for file: %s", + path, + ) + raise def emb_waveform(self, waveform): - waveform = self._prepare_waveform(waveform) + """ + Create an embedding directly from a waveform tensor. + """ + + try: + waveform = self._prepare_waveform(waveform) - with torch.no_grad(): - emb = self.model.encode_batch(waveform)[0].cpu().numpy() + with torch.no_grad(): + embedding = ( + self.model + .encode_batch(waveform) + .squeeze() + .cpu() + .numpy() + ) - return self._normalize(emb) + return self._normalize(embedding) + + except Exception: + logger.exception( + "Embedding extraction failed from waveform." + ) + raise def extract(self, waveform): + """ + Compatibility wrapper. + """ return self.emb_waveform(waveform) \ No newline at end of file diff --git a/src/vocalid/enroll_and_authenticate.py b/src/vocalid/enroll_and_authenticate.py new file mode 100644 index 0000000..88db58e --- /dev/null +++ b/src/vocalid/enroll_and_authenticate.py @@ -0,0 +1,44 @@ +""" +End-to-end demo: + 1. Enroll positive (your voice) and negative (other voices) samples + 2. Train a model on the collected dataset + 3. Authenticate live from the microphone + +Run: python examples/enroll_and_authenticate.py +""" + +import glob + +from vocalid.enrollment import EnrollmentSession +from vocalid.trainer import VoiceTrainer +from vocalid.auth_adapter import VoiceAuthenticator + +DATASET_ROOT = "dataset" +MODEL_PATH = "my_voice_model.pkl" + + +def main(): + session = EnrollmentSession(dataset_root=DATASET_ROOT) + + print("=== Enrolling YOUR voice (positive samples) ===") + session.enroll(label="positive", count=10, seconds=5.0) + + print("\n=== Enrolling OTHER voices (negative samples) ===") + print("Have a few different people speak for this part.") + session.enroll(label="negative", count=10, seconds=5.0) + + print("\n=== Training ===") + pos_files = glob.glob(f"{DATASET_ROOT}/my_voice/*.wav") + neg_files = glob.glob(f"{DATASET_ROOT}/other_voices/*.wav") + trainer = VoiceTrainer() + trainer.train(pos_files, neg_files, save_path=MODEL_PATH) + print(f"Model saved to {MODEL_PATH}") + + print("\n=== Authenticate ===") + authenticator = VoiceAuthenticator(MODEL_PATH) + result = authenticator.authenticate_live(seconds=4.0) + print(result) + + +if __name__ == "__main__": + main() diff --git a/src/vocalid/enrollment.py b/src/vocalid/enrollment.py new file mode 100644 index 0000000..d932b45 --- /dev/null +++ b/src/vocalid/enrollment.py @@ -0,0 +1,187 @@ +""" +enrollment.py + +Ties recorder + validator + dataset_manager +together into one guided session: + + record clip -> validate -> save to dataset + +Retries a rejected clip instead of silently skipping it. +""" +# src/vocalid/typing_compat.py + +from typing import ( + List, + Dict, + Tuple, + Set, + Optional, + Union, + Any, + Callable, +) +import os +import tempfile + +from .recorder import record_one +from .validator import check_audio +from .dataset_manager import DatasetManager +from . import config + + +class EnrollmentSession: + + def __init__( + self, + dataset_root: str = "dataset", + max_attempts_per_sample: int = 3 + ): + + self.dataset = DatasetManager(dataset_root) + self.max_attempts_per_sample = max_attempts_per_sample + + def enroll( + self, + label: str, + count: int = 10, + seconds: float = 5.0 + ) -> List[str]: + + """ + Records and accepts `count` valid clips for + `label` ("positive" or "negative"). + + Returns saved file paths. + + Label options: + + positive: + Target/enrolled speaker voice samples + + negative: + Other speaker voice samples + + Flow: + + record clip + | + validate audio quality + | + save accepted clip + """ + + # Validate label before doing anything + if label not in ("positive", "negative"): + raise ValueError( + "label must be 'positive' or 'negative'" + ) + + print(""" +================================================== + VocalID Enrollment Instructions +================================================== + +Label options: + + positive -> Target/enrolled speaker voice samples + + negative -> Other speaker voice samples + + +Enrollment flow: + + record clip + | + validate audio quality + | + save accepted clip + + +Recording guidelines: + + - Duration: 4-6.5 seconds + - Speak naturally + - Avoid silence + - Avoid background noise + +================================================== +""") + + saved_paths = [] + + for i in range(count): + + accepted = False + + for attempt in range( + 1, + self.max_attempts_per_sample + 1 + ): + + print( + f"\nSample {i + 1}/{count} " + f"(attempt {attempt}) " + f"- label: {label}" + ) + + audio = record_one( + seconds, + config.SAMPLE_RATE + ) + + with tempfile.NamedTemporaryFile( + suffix=".wav", + delete=False + ) as tmp: + + tmp_path = tmp.name + + from .audio_utils import save_audio + + save_audio( + audio, + tmp_path, + config.SAMPLE_RATE + ) + + # Validate audio + is_valid, reason = check_audio( + tmp_path + ) + + if not is_valid: + + print( + f"Rejected: {reason}. Try again." + ) + + os.remove(tmp_path) + continue + + # Save accepted sample + saved_path = self.dataset.add_sample( + tmp_path, + label + ) + + os.remove(tmp_path) + + saved_paths.append( + saved_path + ) + + print( + f"Saved as {saved_path}" + ) + + accepted = True + break + + if not accepted: + + print( + f"Giving up on sample {i + 1} " + f"after {self.max_attempts_per_sample} attempts." + ) + + return saved_paths \ No newline at end of file diff --git a/src/vocalid/recorder.py b/src/vocalid/recorder.py new file mode 100644 index 0000000..e4389ab --- /dev/null +++ b/src/vocalid/recorder.py @@ -0,0 +1,113 @@ +""" +recorder.py + +Records audio clips for enrollment. +""" + + +import os + +from .audio_utils import record_audio, save_audio +from . import config + + + +def record_one( + seconds: float = 5.0, + sample_rate: int = None +): + + """ + Record one audio clip. + """ + + sample_rate = ( + sample_rate + or config.SAMPLE_RATE + ) + + + print( + f"Recording for {seconds:.1f}s... speak naturally." + ) + + + audio = record_audio( + duration=seconds, + sample_rate=sample_rate + ) + + + print( + "Done." + ) + + + return audio + + + + +def record_batch( + out_dir: str, + count: int = 10, + seconds: float = 5.0, + sample_rate: int = None, + prefix: str = "sample" +): + + """ + Record multiple clips and save them. + + Returns: + list of saved paths + """ + + + os.makedirs( + out_dir, + exist_ok=True + ) + + + sample_rate = ( + sample_rate + or config.SAMPLE_RATE + ) + + + paths = [] + + + for i in range(count): + + input( + f"[{i+1}/{count}] Press Enter, then start speaking..." + ) + + + audio = record_one( + seconds, + sample_rate + ) + + + path = os.path.join( + out_dir, + f"{prefix}_{i+1:03d}.wav" + ) + + + save_audio( + audio, + path, + sample_rate + ) + + + paths.append( + path + ) + + + return paths \ No newline at end of file diff --git a/src/vocalid/trainer.py b/src/vocalid/trainer.py index 7dfc6fa..4abe883 100644 --- a/src/vocalid/trainer.py +++ b/src/vocalid/trainer.py @@ -31,7 +31,9 @@ def train(self, positive_paths, negative_paths, save_path="voice_auth.pkl"): y.append(1) for p in negative_paths: + print(f"Loading {p}") emb = self.extractor.embed_file(p) + print(type(emb), None if emb is None else emb.shape) if emb is None: continue X.append(normalize(emb)) diff --git a/src/vocalid/validator.py b/src/vocalid/validator.py new file mode 100644 index 0000000..7ff549b --- /dev/null +++ b/src/vocalid/validator.py @@ -0,0 +1,59 @@ +""" +validator.py +Basic sanity checks on a recorded clip before it's allowed into the +dataset: right length, not silent, not clipped/overdriven. + +Uses audio_utils.load_audio(path) -> (tensor, sample_rate), which the +library already needs for file-based verification. +""" +# src/vocalid/typing_compat.py + +from typing import ( + List, + Dict, + Tuple, + Set, + Optional, + Union, + Any, + Callable, +) +import numpy as np +from .audio_utils import load_audio + +MIN_SECONDS = 4.0 +MAX_SECONDS = 6.5 +MIN_RMS = 0.01 # below this, treat the clip as silence/noise floor +CLIP_THRESHOLD = 0.98 # fraction of full scale considered "clipping" +MAX_CLIPPED_RATIO = 0.01 + + +def _to_numpy(audio): + if hasattr(audio, "numpy"): + return audio.numpy().flatten() + return np.asarray(audio).flatten() + + +def check_audio(path: str) -> Tuple[bool, str]: + """ + Returns (is_valid, reason). reason is 'ok' when valid, otherwise + a short human-readable explanation of why it was rejected. + """ + audio, sample_rate = load_audio(path) + samples = _to_numpy(audio) + + duration = len(samples) / float(sample_rate) + if duration < MIN_SECONDS: + return False, f"too short ({duration:.1f}s)" + if duration > MAX_SECONDS: + return False, f"too long ({duration:.1f}s)" + + rms = float(np.sqrt(np.mean(samples ** 2))) + if rms < MIN_RMS: + return False, "too quiet / silence detected" + + clipped_ratio = float(np.mean(np.abs(samples) >= CLIP_THRESHOLD)) + if clipped_ratio > MAX_CLIPPED_RATIO: + return False, "clipping / distortion detected" + + return True, "ok" diff --git a/src/vocalid/verifier.py b/src/vocalid/verifier.py index 12748c7..4745dc0 100644 --- a/src/vocalid/verifier.py +++ b/src/vocalid/verifier.py @@ -10,7 +10,7 @@ def __init__(self, model_path): self.extractor = EmbeddingExtractor() def verify_file(self, file_path, threshold=THRESHOLD): - wav = load_audio(file_path) # waveform tensor + wav,_ = load_audio(file_path) # waveform tensor emb = self.extractor.emb_waveform(wav) # embedding as numpy score = self.model.predict_proba([emb])[0][1] verified = score >= threshold diff --git a/tests/conftest.py b/tests/conftest.py index 9cfd57c..e133926 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,10 +1,50 @@ -import sys -import os +""" +conftest.py +Shared fixtures for the new modules' tests. Everything here writes +real, tiny .wav files with numpy + soundfile so tests run anywhere +(CI included) without a microphone or a real voice sample. +""" -ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) +import numpy as np +import soundfile as sf +import pytest -if os.environ.get("GITHUB_ACTIONS") == "true": - os.environ["SKIP_SPEECHBRAIN"] = "1" +SAMPLE_RATE = 16000 -if ROOT not in sys.path: - sys.path.insert(0, ROOT) \ No newline at end of file + +def _tone(duration: float, sample_rate: int, freq: float = 220.0, amplitude: float = 0.3): + """A simple sine wave - stands in for a 'real' voice-shaped signal.""" + t = np.linspace(0, duration, int(sample_rate * duration), endpoint=False) + return (amplitude * np.sin(2 * np.pi * freq * t)).astype(np.float32) + + +@pytest.fixture +def write_wav(tmp_path): + """ + Returns a function write_wav(name, **kwargs) -> path + that writes a synthetic clip and returns its path. + + kwargs: + duration seconds (default 5.0) + sample_rate (default 16000) + freq tone frequency, changes the clip's "identity" (default 220.0) + amplitude 0..1 (default 0.3) + silent if True, writes near-zero samples instead of a tone + clipped if True, writes a hard-clipped square-ish wave + """ + def _write(name: str, duration: float = 5.0, sample_rate: int = SAMPLE_RATE, + freq: float = 220.0, amplitude: float = 0.3, + silent: bool = False, clipped: bool = False): + if silent: + samples = np.zeros(int(sample_rate * duration), dtype=np.float32) + elif clipped: + samples = _tone(duration, sample_rate, freq, amplitude=1.0) + samples = np.clip(samples * 5.0, -1.0, 1.0) + else: + samples = _tone(duration, sample_rate, freq, amplitude) + + path = tmp_path / name + sf.write(str(path), samples, sample_rate) + return str(path) + + return _write diff --git a/tests/test_audio_utils.py b/tests/test_audio_utils.py new file mode 100644 index 0000000..7170596 --- /dev/null +++ b/tests/test_audio_utils.py @@ -0,0 +1,204 @@ +import numpy as np +import pytest +import torch + +from unittest.mock import patch + +from vocalid.audio_utils import ( + load_audio, + record_audio, + save_audio, +) + + +# --------------------------------------------------------------------- +# load_audio +# --------------------------------------------------------------------- + +@patch("vocalid.audio_utils.sf.read") +def test_load_audio_mono(mock_read): + + samples = np.random.rand(16000, 1).astype(np.float32) + + mock_read.return_value = ( + samples, + 16000, + ) + + waveform, sr = load_audio("dummy.wav") + + assert sr == 16000 + assert isinstance(waveform, torch.Tensor) + assert waveform.shape == (1, 16000) + + +@patch("vocalid.audio_utils.sf.read") +def test_load_audio_stereo(mock_read): + + samples = np.random.rand(16000, 2).astype(np.float32) + + mock_read.return_value = ( + samples, + 16000, + ) + + waveform, _ = load_audio("dummy.wav") + + assert waveform.shape == (1, 16000) + + +@patch("vocalid.audio_utils.torchaudio.functional.resample") +@patch("vocalid.audio_utils.sf.read") +def test_load_audio_resample( + mock_read, + mock_resample, +): + + samples = np.random.rand(8000, 1).astype(np.float32) + + mock_read.return_value = ( + samples, + 8000, + ) + + mock_resample.return_value = torch.rand(1, 16000) + + waveform, sr = load_audio("dummy.wav") + + assert sr == 16000 + mock_resample.assert_called_once() + + +@patch("vocalid.audio_utils.sf.read") +def test_load_audio_padding(mock_read): + + samples = np.random.rand(4000, 1).astype(np.float32) + + mock_read.return_value = ( + samples, + 16000, + ) + + waveform, _ = load_audio("dummy.wav") + + assert waveform.shape == (1, 16000) + + +@patch("vocalid.audio_utils.sf.read") +def test_load_audio_unsqueeze(mock_read): + """ + Force execution of the waveform.ndim == 1 branch. + """ + + waveform = torch.rand(16000) + + with patch( + "vocalid.audio_utils.torch.from_numpy", + return_value=waveform, + ): + mock_read.return_value = ( + np.random.rand(16000, 1).astype(np.float32), + 16000, + ) + + result, _ = load_audio("dummy.wav") + + assert result.shape == (1, 16000) + + +# --------------------------------------------------------------------- +# record_audio +# --------------------------------------------------------------------- + +@patch("vocalid.audio_utils.sd.wait") +@patch("vocalid.audio_utils.sd.rec") +@patch("vocalid.audio_utils.sd.query_devices") +def test_record_audio_success( + mock_devices, + mock_rec, + mock_wait, +): + + mock_devices.return_value = [ + {"max_input_channels": 1} + ] + + mock_rec.return_value = np.random.rand( + 16000, + 1, + ).astype(np.float32) + + waveform = record_audio(duration=1) + + assert isinstance(waveform, torch.Tensor) + assert waveform.shape == (1, 16000) + + mock_wait.assert_called_once() + + +@patch("vocalid.audio_utils.sd.query_devices") +def test_record_audio_no_microphone( + mock_devices, +): + + mock_devices.return_value = [ + {"max_input_channels": 0} + ] + + with pytest.raises(RuntimeError): + record_audio() + + +@patch("vocalid.audio_utils.sd.rec") +@patch("vocalid.audio_utils.sd.query_devices") +def test_record_audio_recording_failure( + mock_devices, + mock_rec, +): + + mock_devices.return_value = [ + {"max_input_channels": 1} + ] + + mock_rec.side_effect = Exception( + "Recording failed" + ) + + with pytest.raises(RuntimeError): + record_audio() + + +# --------------------------------------------------------------------- +# save_audio +# --------------------------------------------------------------------- + +@patch("vocalid.audio_utils.sf.write") +def test_save_audio(mock_write): + + waveform = torch.rand( + 1, + 16000, + ) + + save_audio( + waveform, + "output.wav", + ) + + mock_write.assert_called_once() + + +@patch("vocalid.audio_utils.sf.write") +def test_save_audio_1d_waveform(mock_write): + """ + Execute the branch where waveform.ndim != 2. + """ + + waveform = torch.rand(16000) + + save_audio( + waveform, + "output.wav", + ) + + mock_write.assert_called_once() \ No newline at end of file diff --git a/tests/test_auth_adapter.py b/tests/test_auth_adapter.py new file mode 100644 index 0000000..6ea7675 --- /dev/null +++ b/tests/test_auth_adapter.py @@ -0,0 +1,103 @@ +""" +test_auth_adapter.py + +VoiceVerifier is mocked so these tests only verify that +VoiceAuthenticator correctly converts verification results +into AuthenticationResult objects. +""" + +import pytest + +import vocalid.auth_adapter as auth_module +from vocalid.auth_adapter import ( + VoiceAuthenticator, + AuthenticationResult, +) + + +class _FakeVerifier: + def __init__(self, model_path): + self.model_path = model_path + + def verify_file(self, path): + return True, 0.87 + + def verify_array(self, audio): + return False, 0.21 + + +@pytest.fixture +def authenticator(monkeypatch): + monkeypatch.setattr( + auth_module, + "VoiceVerifier", + _FakeVerifier + ) + + return VoiceAuthenticator("dummy_model.pkl") + + +def test_authenticate_file_returns_granted_result(authenticator): + + result = authenticator.authenticate_file("some_clip.wav") + + assert isinstance(result, AuthenticationResult) + assert result.success is True + assert result.confidence == 0.87 + + +def test_authenticate_live_returns_denied_result( + authenticator, + monkeypatch +): + + monkeypatch.setattr( + "vocalid.recorder.record_one", + lambda seconds, sample_rate: "fake_audio" + ) + + result = authenticator.authenticate_live( + seconds=4.0 + ) + + assert result.success is False + assert result.confidence == 0.21 + + +def test_authenticate_live_defaults_to_four_seconds( + authenticator, + monkeypatch +): + + captured = {} + + def fake_record_one(seconds, sample_rate): + captured["seconds"] = seconds + return "fake_audio" + + monkeypatch.setattr( + "vocalid.recorder.record_one", + fake_record_one + ) + + authenticator.authenticate_live() + + assert captured["seconds"] == 4.0 + + +def test_result_str_formatting(): + + granted = AuthenticationResult( + success=True, + confidence=0.9123 + ) + + denied = AuthenticationResult( + success=False, + confidence=0.1 + ) + + assert "ACCESS GRANTED" in str(granted) + assert "0.91" in str(granted) + + assert "ACCESS DENIED" in str(denied) \ No newline at end of file diff --git a/tests/test_cli.py b/tests/test_cli.py index 6bca847..149929f 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -1,29 +1,169 @@ + import sys -from unittest.mock import patch, MagicMock +from unittest.mock import MagicMock, patch + from vocalid.cli import main + +# --------------------------------------------------------------------- +# train +# --------------------------------------------------------------------- + @patch("vocalid.cli.VoiceTrainer") def test_cli_train(mock_trainer): - trainer_instance = MagicMock() - mock_trainer.return_value = trainer_instance - test_args = ["cli.py", "train", "--positive", "pos_folder", "--negative", "neg_folder"] - with patch.object(sys, 'argv', test_args): + trainer = MagicMock() + mock_trainer.return_value = trainer + + args = [ + "cli.py", + "train", + "--positive", "pos_folder", + "--negative", "neg_folder", + ] + + with patch.object(sys, "argv", args): + main() + + trainer.train.assert_called_once() + + +# --------------------------------------------------------------------- +# evaluate +# --------------------------------------------------------------------- + +@patch("vocalid.cli.VoiceTrainer") +def test_cli_evaluate(mock_trainer): + + trainer = MagicMock() + + trainer.evaluate.return_value = { + "accuracy": 0.95, + "report": "classification report", + } + + mock_trainer.return_value = trainer + + args = [ + "cli.py", + "evaluate", + "--model", "model.pkl", + "--positive", "pos_folder", + "--negative", "neg_folder", + ] + + with patch.object(sys, "argv", args): main() - # Ensure the train method was called - trainer_instance.train.assert_called_once() + trainer.load.assert_called_once_with("model.pkl") + trainer.evaluate.assert_called_once() + +# --------------------------------------------------------------------- +# verify +# --------------------------------------------------------------------- @patch("vocalid.cli.VoiceVerifier") def test_cli_verify(mock_verifier): - verifier_instance = MagicMock() - verifier_instance.verify_file.return_value = (True, 0.95) - mock_verifier.return_value = verifier_instance - test_args = ["cli.py", "verify", "dummy.wav"] - with patch.object(sys, 'argv', test_args): - main() + verifier = MagicMock() + + verifier.verify_file.return_value = ( + True, + 0.95, + ) + + mock_verifier.return_value = verifier + + args = [ + "cli.py", + "verify", + "dummy.wav", + ] + + with patch.object(sys, "argv", args): + main() + + verifier.verify_file.assert_called_once_with( + "dummy.wav" + ) + + +# --------------------------------------------------------------------- +# live success +# --------------------------------------------------------------------- + +@patch("vocalid.cli.record_audio") +@patch("vocalid.cli.VoiceVerifier") +def test_cli_live_success( + mock_verifier, + mock_record, +): + + verifier = MagicMock() + + verifier.verify_array.return_value = ( + True, + 0.91, + ) + + mock_verifier.return_value = verifier + + mock_record.return_value = MagicMock() + + args = [ + "cli.py", + "live", + "--seconds", + "5", + ] + + with patch.object(sys, "argv", args): + main() + + mock_record.assert_called_once_with(5) + + verifier.verify_array.assert_called_once() + + +# --------------------------------------------------------------------- +# live microphone failure +# --------------------------------------------------------------------- + +@patch("vocalid.cli.record_audio") +@patch("vocalid.cli.VoiceVerifier") +def test_cli_live_microphone_failure( + mock_verifier, + mock_record, +): + + mock_verifier.return_value = MagicMock() + + mock_record.side_effect = RuntimeError( + "No microphone" + ) + + args = [ + "cli.py", + "live", + ] + + with patch.object(sys, "argv", args): + main() + + +# --------------------------------------------------------------------- +# help +# --------------------------------------------------------------------- + +@patch("argparse.ArgumentParser.print_help") +def test_cli_print_help(mock_help): + + args = [ + "cli.py", + ] + + with patch.object(sys, "argv", args): + main() - verifier_instance.verify_file.assert_called_once_with("dummy.wav") - \ No newline at end of file + mock_help.assert_called_once() \ No newline at end of file diff --git a/tests/test_dataset_manager.py b/tests/test_dataset_manager.py new file mode 100644 index 0000000..ebf4b54 --- /dev/null +++ b/tests/test_dataset_manager.py @@ -0,0 +1,46 @@ +""" +test_dataset_manager.py +""" + +import os +import pytest +from vocalid.dataset_manager import DatasetManager + + +def test_creates_folder_layout(tmp_path): + root = tmp_path / "dataset" + DatasetManager(str(root)) + assert (root / "my_voice").is_dir() + assert (root / "other_voices").is_dir() + + +def test_add_sample_copies_and_lists(tmp_path, write_wav): + dm = DatasetManager(str(tmp_path / "dataset")) + clip = write_wav("clip.wav") + + saved_path = dm.add_sample(clip, "positive") + + assert os.path.exists(saved_path) + assert saved_path in dm.list_samples("positive") + assert dm.list_samples("negative") == [] + + +def test_add_sample_sequential_naming(tmp_path, write_wav): + dm = DatasetManager(str(tmp_path / "dataset")) + clip_a = write_wav("a.wav") + clip_b = write_wav("b.wav") + + path1 = dm.add_sample(clip_a, "negative") + path2 = dm.add_sample(clip_b, "negative") + + assert os.path.basename(path1) == "sample001.wav" + assert os.path.basename(path2) == "sample002.wav" + assert len(dm.list_samples("negative")) == 2 + + +def test_rejects_bad_label(tmp_path, write_wav): + dm = DatasetManager(str(tmp_path / "dataset")) + clip = write_wav("clip.wav") + + with pytest.raises(ValueError): + dm.add_sample(clip, "unknown_label") diff --git a/tests/test_embeddings.py b/tests/test_embeddings.py index 49d8f98..31f0e59 100644 --- a/tests/test_embeddings.py +++ b/tests/test_embeddings.py @@ -1,26 +1,309 @@ import numpy as np -from unittest.mock import patch, MagicMock import pytest +import torch -@patch("vocalid.embeddings.EmbeddingExtractor") -def test_embedding_extractor(mock_extractor): +from unittest.mock import MagicMock, patch + +from vocalid.embeddings import EmbeddingExtractor + + +def make_extractor(): """ - Ensure that the EmbeddingExtractor returns a fixed-size embedding - without requiring SpeechBrain to be installed. + Create an EmbeddingExtractor without loading SpeechBrain. """ - # Mock the extractor instance - mock_instance = MagicMock() - mock_instance.extract.return_value = np.random.rand(192) # fixed embedding size - mock_extractor.return_value = mock_instance + extractor = EmbeddingExtractor.__new__(EmbeddingExtractor) + extractor.model = MagicMock() + return extractor + + +# --------------------------------------------------------------------- +# __init__ +# --------------------------------------------------------------------- + +def test_embedding_extractor_init(): + + mock_encoder = MagicMock() + mock_model = MagicMock() + + mock_encoder.from_hparams.return_value = mock_model + + fake_speechbrain = MagicMock() + fake_speechbrain.inference.EncoderClassifier = mock_encoder + + with patch.dict( + "sys.modules", + { + "speechbrain": fake_speechbrain, + "speechbrain.inference": fake_speechbrain.inference, + }, + ): + + extractor = EmbeddingExtractor() + + assert extractor.model == mock_model + + mock_encoder.from_hparams.assert_called_once_with( + source="speechbrain/spkrec-ecapa-voxceleb", + run_opts={"device": "cpu"}, + savedir="pretrained_models/ecapa", + ) + + +def test_embedding_extractor_import_error(): + + with patch.dict( + "sys.modules", + { + "speechbrain": None, + "speechbrain.inference": None, + }, + ): + + with pytest.raises(ImportError) as exc: + + EmbeddingExtractor() + + assert "SpeechBrain could not load" in str(exc.value) + + +# --------------------------------------------------------------------- +# _normalize +# --------------------------------------------------------------------- + +def test_normalize_vector(): + + extractor = make_extractor() + + vec = np.array([3.0, 4.0]) + + result = extractor._normalize(vec) + + assert np.allclose( + result, + np.array([0.6, 0.8]) + ) + + +def test_normalize_zero_vector(): + + extractor = make_extractor() + + vec = np.zeros(5) + + result = extractor._normalize(vec) + + assert np.array_equal( + result, + vec + ) + + +# --------------------------------------------------------------------- +# _prepare_waveform +# --------------------------------------------------------------------- + +def test_prepare_waveform_numpy(): + + extractor = make_extractor() + + waveform = np.random.rand( + 16000 + ).astype(np.float32) + + result = extractor._prepare_waveform( + waveform + ) + + assert isinstance( + result, + torch.Tensor + ) + + assert result.shape == ( + 1, + 16000 + ) + + +def test_prepare_waveform_stereo(): + + extractor = make_extractor() + + waveform = torch.rand( + 2, + 16000 + ) + + result = extractor._prepare_waveform( + waveform + ) + + assert result.shape == ( + 1, + 16000 + ) + + +def test_prepare_waveform_padding(): + + extractor = make_extractor() + + waveform = np.random.rand( + 8000 + ).astype(np.float32) + + result = extractor._prepare_waveform( + waveform + ) + + assert result.shape == ( + 1, + 16000 + ) + + +# --------------------------------------------------------------------- +# embed_file +# --------------------------------------------------------------------- + +@patch("vocalid.embeddings.load_audio") +def test_embed_file_success(mock_load_audio): + + extractor = make_extractor() + + waveform = torch.rand( + 1, + 16000 + ) + + mock_load_audio.return_value = ( + waveform, + 16000 + ) + + extractor.model.encode_batch.return_value = torch.rand( + 1, + 192 + ) + + result = extractor.embed_file( + "dummy.wav" + ) + + assert isinstance( + result, + np.ndarray + ) + + assert result.shape == ( + 192, + ) + + mock_load_audio.assert_called_once_with( + "dummy.wav" + ) + + +@patch("vocalid.embeddings.load_audio") +def test_embed_file_exception(mock_load_audio): + + extractor = make_extractor() + + mock_load_audio.side_effect = RuntimeError( + "Load failed" + ) + + with pytest.raises(RuntimeError): + + extractor.embed_file( + "dummy.wav" + ) + + +# --------------------------------------------------------------------- +# emb_waveform +# --------------------------------------------------------------------- + +def test_emb_waveform_success(): + + extractor = make_extractor() + + waveform = torch.rand( + 1, + 16000 + ) + + extractor.model.encode_batch.return_value = torch.rand( + 1, + 192 + ) + + result = extractor.emb_waveform( + waveform + ) + + assert isinstance( + result, + np.ndarray + ) + + assert result.shape == ( + 192, + ) + + +def test_emb_waveform_exception(): + + extractor = make_extractor() + + waveform = torch.rand( + 1, + 16000 + ) + + extractor.model.encode_batch.side_effect = RuntimeError( + "Encoding failed" + ) + + with pytest.raises(RuntimeError): + + extractor.emb_waveform( + waveform + ) + + +# --------------------------------------------------------------------- +# extract +# --------------------------------------------------------------------- + +def test_extract_calls_emb_waveform(): + + extractor = make_extractor() + + waveform = torch.rand( + 1, + 16000 + ) + + expected = np.random.rand( + 192 + ) - # Import inside the test to ensure mock patching works - from vocalid.embeddings import EmbeddingExtractor + with patch.object( + extractor, + "emb_waveform", + return_value=expected, + ) as mock_method: - extractor = EmbeddingExtractor() - dummy_input = np.random.rand(16000).astype("float32") # 1-second dummy audio at 16kHz + result = extractor.extract( + waveform + ) - emb = extractor.extract(dummy_input) + mock_method.assert_called_once_with( + waveform + ) - # Check the embedding shape - assert emb.shape[0] == 192 - mock_instance.extract.assert_called_once_with(dummy_input) + assert np.array_equal( + result, + expected + ) \ No newline at end of file diff --git a/tests/test_enroll_and_authenticate.py b/tests/test_enroll_and_authenticate.py new file mode 100644 index 0000000..8cc9c25 --- /dev/null +++ b/tests/test_enroll_and_authenticate.py @@ -0,0 +1,120 @@ +from unittest.mock import MagicMock, patch + +from vocalid.enroll_and_authenticate import ( + main, + DATASET_ROOT, + MODEL_PATH, +) + + +@patch("vocalid.enroll_and_authenticate.print") +@patch("vocalid.enroll_and_authenticate.glob.glob") +@patch("vocalid.enroll_and_authenticate.VoiceAuthenticator") +@patch("vocalid.enroll_and_authenticate.VoiceTrainer") +@patch("vocalid.enroll_and_authenticate.EnrollmentSession") +def test_main( + mock_session_cls, + mock_trainer_cls, + mock_auth_cls, + mock_glob, + mock_print, +): + """ + Test the complete enroll -> train -> authenticate workflow. + """ + + # ----------------------------- + # Enrollment session + # ----------------------------- + session = MagicMock() + mock_session_cls.return_value = session + + # ----------------------------- + # Trainer + # ----------------------------- + trainer = MagicMock() + mock_trainer_cls.return_value = trainer + + # ----------------------------- + # Authenticator + # ----------------------------- + authenticator = MagicMock() + authenticator.authenticate_live.return_value = { + "verified": True, + "score": 0.96, + } + mock_auth_cls.return_value = authenticator + + # ----------------------------- + # Fake dataset files + # ----------------------------- + mock_glob.side_effect = [ + [ + f"{DATASET_ROOT}/my_voice/a.wav", + f"{DATASET_ROOT}/my_voice/b.wav", + ], + [ + f"{DATASET_ROOT}/other_voices/x.wav", + f"{DATASET_ROOT}/other_voices/y.wav", + ], + ] + + # ----------------------------- + # Run + # ----------------------------- + main() + + # ----------------------------- + # Enrollment + # ----------------------------- + mock_session_cls.assert_called_once_with( + dataset_root=DATASET_ROOT + ) + + session.enroll.assert_any_call( + label="positive", + count=10, + seconds=5.0, + ) + + session.enroll.assert_any_call( + label="negative", + count=10, + seconds=5.0, + ) + + assert session.enroll.call_count == 2 + + # ----------------------------- + # Dataset lookup + # ----------------------------- + assert mock_glob.call_count == 2 + + # ----------------------------- + # Training + # ----------------------------- + trainer.train.assert_called_once_with( + [ + f"{DATASET_ROOT}/my_voice/a.wav", + f"{DATASET_ROOT}/my_voice/b.wav", + ], + [ + f"{DATASET_ROOT}/other_voices/x.wav", + f"{DATASET_ROOT}/other_voices/y.wav", + ], + save_path=MODEL_PATH, + ) + + # ----------------------------- + # Authentication + # ----------------------------- + mock_auth_cls.assert_called_once_with(MODEL_PATH) + + authenticator.authenticate_live.assert_called_once_with( + seconds=4.0 + ) + + # ----------------------------- + # Output + # ----------------------------- + assert mock_print.call_count >= 6 \ No newline at end of file diff --git a/tests/test_enrollment.py b/tests/test_enrollment.py new file mode 100644 index 0000000..653ec85 --- /dev/null +++ b/tests/test_enrollment.py @@ -0,0 +1,151 @@ +""" +test_enrollment.py + +Unit tests for EnrollmentSession. + +External dependencies are mocked: + +- recorder.record_one +- validator.check_audio +- audio_utils.save_audio + +DatasetManager is used normally against pytest's temporary directory. +""" + +import pytest + +import vocalid.enrollment as enrollment_module +from vocalid.enrollment import EnrollmentSession + + +# --------------------------------------------------------- +# Shared fixtures +# --------------------------------------------------------- + +@pytest.fixture(autouse=True) +def patch_recording(monkeypatch): + + monkeypatch.setattr( + enrollment_module, + "record_one", + lambda seconds, sample_rate: "fake_audio" + ) + + monkeypatch.setattr( + "vocalid.audio_utils.save_audio", + lambda audio, path, sr: None + ) + + +# --------------------------------------------------------- +# Tests +# --------------------------------------------------------- + +def test_accepts_valid_clip(monkeypatch, tmp_path): + + monkeypatch.setattr( + enrollment_module, + "check_audio", + lambda path: (True, "ok") + ) + + session = EnrollmentSession( + dataset_root=str(tmp_path / "dataset") + ) + + saved = session.enroll( + label="positive", + count=1, + seconds=5.0 + ) + + assert len(saved) == 1 + + +def test_retries_after_validation_failure(monkeypatch, tmp_path): + + results = iter([ + (False, "too quiet"), + (True, "ok") + ]) + + monkeypatch.setattr( + enrollment_module, + "check_audio", + lambda path: next(results) + ) + + session = EnrollmentSession( + dataset_root=str(tmp_path / "dataset"), + max_attempts_per_sample=3 + ) + + saved = session.enroll( + label="positive", + count=1, + seconds=5.0 + ) + + assert len(saved) == 1 + + +def test_gives_up_after_max_attempts(monkeypatch, tmp_path): + + monkeypatch.setattr( + enrollment_module, + "check_audio", + lambda path: (False, "too quiet") + ) + + session = EnrollmentSession( + dataset_root=str(tmp_path / "dataset"), + max_attempts_per_sample=2 + ) + + saved = session.enroll( + label="positive", + count=1, + seconds=5.0 + ) + + assert saved == [] + + +def test_returns_path_for_every_sample(monkeypatch, tmp_path): + + monkeypatch.setattr( + enrollment_module, + "check_audio", + lambda path: (True, "ok") + ) + + session = EnrollmentSession( + dataset_root=str(tmp_path / "dataset") + ) + + saved = session.enroll( + label="negative", + count=3, + seconds=5.0 + ) + + assert len(saved) == 3 + assert len(set(saved)) == 3 + + +def test_invalid_label_is_rejected(tmp_path): + + session = EnrollmentSession( + dataset_root=str(tmp_path / "dataset") + ) + + with pytest.raises( + ValueError, + match="label must be 'positive' or 'negative'" + ): + + session.enroll( + label="kai", + count=1, + seconds=5.0 + ) \ No newline at end of file diff --git a/tests/test_model_store.py b/tests/test_model_store.py new file mode 100644 index 0000000..3cdbf8a --- /dev/null +++ b/tests/test_model_store.py @@ -0,0 +1,45 @@ +from unittest.mock import MagicMock, patch + +from vocalid.model_store import save_model, load_model + + +# --------------------------------------------------------------------- +# save_model +# --------------------------------------------------------------------- + +@patch("vocalid.model_store.joblib.dump") +def test_save_model(mock_dump): + + model = MagicMock() + + save_model( + model, + "dummy.pkl", + ) + + mock_dump.assert_called_once_with( + model, + "dummy.pkl", + ) + + +# --------------------------------------------------------------------- +# load_model +# --------------------------------------------------------------------- + +@patch("vocalid.model_store.joblib.load") +def test_load_model(mock_load): + + fake_model = MagicMock() + + mock_load.return_value = fake_model + + model = load_model( + "dummy.pkl", + ) + + mock_load.assert_called_once_with( + "dummy.pkl", + ) + + assert model == fake_model \ No newline at end of file diff --git a/tests/test_recorder.py b/tests/test_recorder.py new file mode 100644 index 0000000..28a528c --- /dev/null +++ b/tests/test_recorder.py @@ -0,0 +1,66 @@ +""" +test_recorder.py + +record_audio / save_audio (from audio_utils) are mocked out - these +tests check recorder.py's own logic (sample-rate defaulting, file +naming, loop count), not the real microphone or real file writing. +""" + +import pytest +import vocalid.recorder as recorder_module + + +@pytest.fixture +def fake_audio_utils(monkeypatch): + calls = {"record_audio": [], "save_audio": []} + + def fake_record_audio(duration, sample_rate): + calls["record_audio"].append((duration, sample_rate)) + return f"audio@{duration}s" + + def fake_save_audio(audio, path, sample_rate): + calls["save_audio"].append((audio, path, sample_rate)) + + monkeypatch.setattr(recorder_module, "record_audio", fake_record_audio) + monkeypatch.setattr(recorder_module, "save_audio", fake_save_audio) + return calls + + +def test_record_one_uses_given_seconds_and_rate(fake_audio_utils): + audio = recorder_module.record_one(seconds=6.0, sample_rate=22050) + assert audio == "audio@6.0s" + assert fake_audio_utils["record_audio"] == [(6.0, 22050)] + + +def test_record_one_falls_back_to_config_sample_rate(fake_audio_utils, monkeypatch): + monkeypatch.setattr(recorder_module.config, "SAMPLE_RATE", 16000) + recorder_module.record_one(seconds=5.0, sample_rate=None) + assert fake_audio_utils["record_audio"] == [(5.0, 16000)] + + +def test_record_batch_records_and_saves_expected_count(fake_audio_utils, monkeypatch, tmp_path): + monkeypatch.setattr("builtins.input", lambda prompt="": "") + + paths = recorder_module.record_batch(str(tmp_path), count=3, seconds=5.0, sample_rate=16000) + + assert len(paths) == 3 + assert len(fake_audio_utils["record_audio"]) == 3 + assert len(fake_audio_utils["save_audio"]) == 3 + + +def test_record_batch_names_files_sequentially(fake_audio_utils, monkeypatch, tmp_path): + monkeypatch.setattr("builtins.input", lambda prompt="": "") + + paths = recorder_module.record_batch(str(tmp_path), count=2, prefix="clip") + + assert paths[0].endswith("clip_001.wav") + assert paths[1].endswith("clip_002.wav") + + +def test_record_batch_creates_output_dir(fake_audio_utils, monkeypatch, tmp_path): + monkeypatch.setattr("builtins.input", lambda prompt="": "") + out_dir = tmp_path / "nested" / "clips" + + recorder_module.record_batch(str(out_dir), count=1) + + assert out_dir.is_dir() diff --git a/tests/test_trainer.py b/tests/test_trainer.py index 3ae682f..12f1555 100644 --- a/tests/test_trainer.py +++ b/tests/test_trainer.py @@ -1,58 +1,268 @@ + + import numpy as np -from unittest.mock import patch, MagicMock -from vocalid.trainer import VoiceTrainer +import pytest + +from unittest.mock import MagicMock, patch + +from vocalid.trainer import ( + VoiceTrainer, + normalize, +) + + +# --------------------------------------------------------------------- +# normalize +# --------------------------------------------------------------------- + +def test_normalize_vector(): + + vec = np.array([3.0, 4.0]) + + result = normalize(vec) + + assert np.allclose( + result, + np.array([0.6, 0.8]), + ) + + +def test_normalize_zero_vector(): + + vec = np.zeros(5) + + result = normalize(vec) + + assert np.array_equal(result, vec) + + +# --------------------------------------------------------------------- +# helper +# --------------------------------------------------------------------- + +def make_trainer(): + + trainer = VoiceTrainer.__new__(VoiceTrainer) + + trainer.extractor = MagicMock() + + trainer.model = None + + return trainer + -# ------------------ Test training pipeline ------------------ # +# --------------------------------------------------------------------- +# train +# --------------------------------------------------------------------- + +@patch("vocalid.trainer.save_model") +@patch("vocalid.trainer.LogisticRegression") @patch("vocalid.trainer.EmbeddingExtractor") -def test_training_pipeline(mock_extractor): - mock_instance = MagicMock() - # embed_file returns 192-dim embeddings - mock_instance.embed_file.side_effect = lambda x: np.random.rand(192) - mock_extractor.return_value = mock_instance +def test_training_pipeline( + mock_extractor, + mock_lr, + mock_save, +): + + extractor = MagicMock() + + extractor.embed_file.side_effect = ( + lambda _: np.random.rand(192) + ) - X_positive = ["file1.wav", "file2.wav", "file3.wav", "file4.wav", "file5.wav"] - X_negative = ["file6.wav", "file7.wav", "file8.wav", "file9.wav", "file10.wav"] + mock_extractor.return_value = extractor + + model = MagicMock() + + mock_lr.return_value = model trainer = VoiceTrainer() - trainer.train(X_positive, X_negative, save_path="dummy_model.pkl") - # Predict using embeddings - # X_test = X_positive + X_negative - # X_features = np.vstack([trainer.extractor.embed_file(x) for x in X_test]) - # y_test = np.array([1]*5 + [0]*5) - # preds = trainer.model.predict(X_features) + positive = [ + "p1.wav", + "p2.wav", + ] + + negative = [ + "n1.wav", + "n2.wav", + ] + + path = trainer.train( + positive, + negative, + "dummy.pkl", + ) + + assert path == "dummy.pkl" + + model.fit.assert_called_once() + + mock_save.assert_called_once() + + assert trainer.model == model - # assert (preds == y_test).mean() >= 0.5 - assert trainer.model is not None +@patch("vocalid.trainer.save_model") +@patch("vocalid.trainer.LogisticRegression") +@patch("vocalid.trainer.EmbeddingExtractor") +def test_train_skips_none_embeddings( + mock_extractor, + mock_lr, + mock_save, +): + + extractor = MagicMock() + + extractor.embed_file.side_effect = [ + np.random.rand(192), + None, + np.random.rand(192), + ] + + mock_extractor.return_value = extractor + + mock_lr.return_value = MagicMock() + + trainer = VoiceTrainer() + + trainer.train( + ["a.wav", "b.wav"], + ["c.wav"], + ) + + trainer.model.fit.assert_called_once() -# ------------------ Test save/load ------------------ # @patch("vocalid.trainer.EmbeddingExtractor") -def test_save_and_load(mock_extractor, tmp_path): - mock_instance = MagicMock() - mock_instance.embed_file.side_effect = lambda x: np.random.rand(192) - mock_extractor.return_value = mock_instance +def test_train_no_embeddings( + mock_extractor, +): + + extractor = MagicMock() + + extractor.embed_file.return_value = None - X_positive = ["file1.wav", "file2.wav", "file3.wav", "file4.wav", "file5.wav"] - X_negative = ["file6.wav", "file7.wav", "file8.wav", "file9.wav", "file10.wav"] + mock_extractor.return_value = extractor trainer = VoiceTrainer() - trainer.train(X_positive, X_negative, save_path=str(tmp_path / "dummy_model.pkl")) - # Save/load test - save_path = tmp_path / "clf.pkl" - trainer.save(str(save_path)) - trainer2 = VoiceTrainer() - trainer2.load(str(save_path)) + with pytest.raises(ValueError): + + trainer.train( + ["a.wav"], + ["b.wav"], + ) + + +# --------------------------------------------------------------------- +# evaluate +# --------------------------------------------------------------------- + +def test_evaluate_success(): + + trainer = make_trainer() + + trainer.model = MagicMock() + + trainer.model.predict.return_value = np.array([1, 0]) + + trainer.extractor.embed_file.side_effect = [ + np.random.rand(192), + np.random.rand(192), + ] + + result = trainer.evaluate( + ["p.wav"], + ["n.wav"], + ) + + assert "accuracy" in result + + assert "report" in result + + +def test_evaluate_without_model(): + + trainer = make_trainer() + + with pytest.raises(ValueError): + + trainer.evaluate( + ["p.wav"], + ["n.wav"], + ) + + +def test_evaluate_empty_embeddings(): + + trainer = make_trainer() + + trainer.model = MagicMock() + + trainer.extractor.embed_file.return_value = None + + with pytest.raises(ValueError): + + trainer.evaluate( + ["p.wav"], + ["n.wav"], + ) + + +# --------------------------------------------------------------------- +# save +# --------------------------------------------------------------------- + +@patch("vocalid.trainer.save_model") +def test_save_success(mock_save): + + trainer = make_trainer() + + trainer.model = MagicMock() + + trainer.save("abc.pkl") + + mock_save.assert_called_once_with( + trainer.model, + "abc.pkl", + ) + + +def test_save_without_model(): + + trainer = make_trainer() + + with pytest.raises(ValueError): + + trainer.save("abc.pkl") + + +# --------------------------------------------------------------------- +# load +# --------------------------------------------------------------------- + +@patch("vocalid.trainer.load_model") +def test_load_success(mock_load): + + trainer = make_trainer() + + model = MagicMock() + + mock_load.return_value = model + + trainer.load("abc.pkl") + + assert trainer.model == model + + +@patch("vocalid.trainer.load_model") +def test_load_failure(mock_load): - # X_test = X_positive + X_negative - # X_features = np.vstack([trainer.extractor.embed_file(x) for x in X_test]) + trainer = make_trainer() - # preds_original = trainer.model.predict(X_features) - # preds_loaded = trainer2.model.predict(X_features) + mock_load.return_value = None - # assert np.allclose(preds_original, preds_loaded) + with pytest.raises(ValueError): - assert trainer.model is not None - assert trainer2.model is not None \ No newline at end of file + trainer.load("abc.pkl") \ No newline at end of file diff --git a/tests/test_validator.py b/tests/test_validator.py new file mode 100644 index 0000000..24f95c9 --- /dev/null +++ b/tests/test_validator.py @@ -0,0 +1,40 @@ +""" +test_validator.py +""" + +from vocalid.validator import check_audio + + +def test_accepts_good_clip(write_wav): + path = write_wav("good.wav", duration=5.0) + is_valid, reason = check_audio(path) + assert is_valid is True + assert reason == "ok" + + +def test_rejects_too_short(write_wav): + path = write_wav("short.wav", duration=1.5) + is_valid, reason = check_audio(path) + assert is_valid is False + assert "short" in reason + + +def test_rejects_too_long(write_wav): + path = write_wav("long.wav", duration=8.0) + is_valid, reason = check_audio(path) + assert is_valid is False + assert "long" in reason + + +def test_rejects_silence(write_wav): + path = write_wav("silent.wav", duration=5.0, silent=True) + is_valid, reason = check_audio(path) + assert is_valid is False + assert "quiet" in reason or "silence" in reason + + +def test_rejects_clipping(write_wav): + path = write_wav("clipped.wav", duration=5.0, clipped=True) + is_valid, reason = check_audio(path) + assert is_valid is False + assert "clip" in reason diff --git a/tests/test_verifier.py b/tests/test_verifier.py index 0277d6d..285f948 100644 --- a/tests/test_verifier.py +++ b/tests/test_verifier.py @@ -1,28 +1,202 @@ -from unittest.mock import patch, MagicMock + + import numpy as np +import torch + +from unittest.mock import MagicMock, patch + +from vocalid.verifier import VoiceVerifier +from vocalid.config import THRESHOLD + + +def make_verifier(): + """ + Create a VoiceVerifier without calling __init__. + """ + verifier = VoiceVerifier.__new__(VoiceVerifier) + verifier.model = MagicMock() + verifier.extractor = MagicMock() + return verifier + + +# --------------------------------------------------------------------- +# __init__ +# --------------------------------------------------------------------- -@patch("vocalid.verifier.load_model") @patch("vocalid.verifier.EmbeddingExtractor") -def test_verifier(mock_extractor, mock_load_model): - # Mock extractor - mock_instance = MagicMock() - mock_instance.extract.return_value = np.zeros(192) - mock_extractor.return_value = mock_instance - - # Mock model - mock_model = MagicMock() - mock_model.predict_proba.return_value = [[0.1, 0.9]] - mock_load_model.return_value = mock_model - - from vocalid.verifier import VoiceVerifier - verifier = VoiceVerifier("model.bin") - - # Fix: call extractor.extract instead of get_embedding - dummy_waveform = np.random.rand(16000).astype("float32") - emb = verifier.extractor.extract(dummy_waveform) - score = mock_model.predict_proba([emb])[0][1] - verified = score >= 0.5 # use some threshold - - assert isinstance(verified, bool) - assert isinstance(score, float) - mock_instance.extract.assert_called_once_with(dummy_waveform) +@patch("vocalid.verifier.load_model") +def test_verifier_init( + mock_load_model, + mock_extractor, +): + fake_model = MagicMock() + fake_extractor = MagicMock() + + mock_load_model.return_value = fake_model + mock_extractor.return_value = fake_extractor + + verifier = VoiceVerifier("dummy_model.pkl") + + mock_load_model.assert_called_once_with( + "dummy_model.pkl" + ) + + mock_extractor.assert_called_once() + + assert verifier.model is fake_model + assert verifier.extractor is fake_extractor + + +# --------------------------------------------------------------------- +# verify_file +# --------------------------------------------------------------------- + +@patch("vocalid.verifier.load_audio") +def test_verify_file_success(mock_load_audio): + + verifier = make_verifier() + + waveform = torch.rand(1, 16000) + + mock_load_audio.return_value = ( + waveform, + 16000, + ) + + verifier.extractor.emb_waveform.return_value = np.random.rand(192) + + verifier.model.predict_proba.return_value = [ + [0.1, 0.9] + ] + + verified, score = verifier.verify_file( + "dummy.wav" + ) + + assert verified is True + assert score == 0.9 + + verifier.extractor.emb_waveform.assert_called_once_with( + waveform + ) + + +@patch("vocalid.verifier.load_audio") +def test_verify_file_threshold(mock_load_audio): + + verifier = make_verifier() + + waveform = torch.rand(1, 16000) + + mock_load_audio.return_value = ( + waveform, + 16000, + ) + + verifier.extractor.emb_waveform.return_value = np.random.rand(192) + + verifier.model.predict_proba.return_value = [ + [0.8, 0.2] + ] + + verified, score = verifier.verify_file( + "dummy.wav" + ) + + assert verified is False + assert score == 0.2 + + +# --------------------------------------------------------------------- +# verify_array +# --------------------------------------------------------------------- + +def test_verify_array_numpy(): + + verifier = make_verifier() + + verifier.extractor.emb_waveform.return_value = np.random.rand(192) + + verifier.model.predict_proba.return_value = [ + [0.2, 0.8] + ] + + waveform = np.random.rand( + 16000 + ).astype(np.float32) + + verified, score = verifier.verify_array( + waveform + ) + + assert verified is True + assert score == 0.8 + + verifier.extractor.emb_waveform.assert_called_once() + + +def test_verify_array_tensor(): + + verifier = make_verifier() + + verifier.extractor.emb_waveform.return_value = np.random.rand(192) + + verifier.model.predict_proba.return_value = [ + [0.6, 0.4] + ] + + waveform = torch.rand( + 1, + 16000, + ) + + verified, score = verifier.verify_array( + waveform + ) + + assert verified is False + assert score == 0.4 + + verifier.extractor.emb_waveform.assert_called_once_with( + waveform + ) + + +# --------------------------------------------------------------------- +# verify_live +# --------------------------------------------------------------------- + +@patch("vocalid.verifier.record_audio") +def test_verify_live(mock_record): + + verifier = make_verifier() + + waveform = torch.rand( + 1, + 16000, + ) + + mock_record.return_value = waveform + + verifier.verify_array = MagicMock( + return_value=( + True, + 0.95, + ) + ) + + verified, score = verifier.verify_live( + seconds=5 + ) + + mock_record.assert_called_once_with( + 5 + ) + + verifier.verify_array.assert_called_once_with( + waveform, + THRESHOLD, + ) + + assert verified is True + assert score == 0.95 \ No newline at end of file