Instructions to use deepsafe/deepsafe-services with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use deepsafe/deepsafe-services with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("deepsafe/deepsafe-services", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download audio/shiftyspeech/tests/test_api.py from deepsafe/deepsafe-services: direct link, hf CLI and curl.
- Browser
- Download file 8.82 kB
-
https://huggingface.co/deepsafe/deepsafe-services/resolve/main/audio/shiftyspeech/tests/test_api.py
- Command line
-
hf download hf://deepsafe/deepsafe-services/audio/shiftyspeech/tests/test_api.py
-
curl -L -o test_api.py https://huggingface.co/deepsafe/deepsafe-services/resolve/main/audio/shiftyspeech/tests/test_api.py
8.82 kB
| """Tests for ShiftySpeech Audio Deepfake Detection API. | |
| Tests cover: | |
| - Health endpoint | |
| - Predict endpoint with real audio | |
| - Predict endpoint with fake audio | |
| - Audio preprocessing (pad/trim/tile) | |
| - Error handling for invalid input | |
| - Response schema validation | |
| """ | |
| import base64 | |
| import io | |
| import os | |
| import sys | |
| import warnings | |
| import numpy as np | |
| import pytest | |
| warnings.filterwarnings("ignore", category=DeprecationWarning) | |
| # Monkey-patch omegaconf before importing api module | |
| import omegaconf._utils as _omegaconf_utils | |
| if not hasattr(_omegaconf_utils, "is_primitive_type"): | |
| _omegaconf_utils.is_primitive_type = lambda t: t in (int, float, bool, str, bytes) | |
| # Add model code path for local testing | |
| SERVICE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| MODEL_CODE_PATH = os.path.join( | |
| SERVICE_DIR, "synthetic_speech_detection", "SSL_Anti-spoofing" | |
| ) | |
| if MODEL_CODE_PATH not in sys.path: | |
| sys.path.insert(0, MODEL_CODE_PATH) | |
| # Patch the api module constants for local testing | |
| import api | |
| api.MODEL_CODE_PATH = MODEL_CODE_PATH | |
| api.WEIGHTS_PATH = os.path.join(SERVICE_DIR, "weights", "hfg_aug_1_2.pt") | |
| api.XLSR_DIR = os.path.join(SERVICE_DIR, "models") | |
| from fastapi.testclient import TestClient | |
| client = TestClient(api.app) | |
| # Dataset paths | |
| DATASET_DIR = os.path.join( | |
| os.path.dirname(SERVICE_DIR), | |
| os.pardir, | |
| os.pardir, | |
| "dataset", | |
| "audio", | |
| ) | |
| DATASET_DIR = os.path.normpath(DATASET_DIR) | |
| REAL_DIR = os.path.join(DATASET_DIR, "real") | |
| FAKE_DIR = os.path.join(DATASET_DIR, "fake") | |
| # Check if model weights are available for integration tests | |
| WEIGHTS_AVAILABLE = os.path.exists(api.WEIGHTS_PATH) and os.path.exists( | |
| os.path.join(SERVICE_DIR, "models", "xlsr2_300m.pt") | |
| ) | |
| DATASET_AVAILABLE = os.path.isdir(REAL_DIR) and os.path.isdir(FAKE_DIR) | |
| def _make_wav_bytes(duration_s: float = 1.0, sr: int = 16000) -> bytes: | |
| """Generate a simple sine wave WAV file as bytes.""" | |
| import soundfile as sf | |
| t = np.linspace(0, duration_s, int(sr * duration_s), endpoint=False) | |
| audio = 0.5 * np.sin(2 * np.pi * 440 * t).astype(np.float32) | |
| buf = io.BytesIO() | |
| sf.write(buf, audio, sr, format="WAV") | |
| buf.seek(0) | |
| return buf.read() | |
| def _encode_file(path: str) -> str: | |
| """Read a file and return base64 encoded string.""" | |
| with open(path, "rb") as f: | |
| return base64.b64encode(f.read()).decode("utf-8") | |
| class TestHealthEndpoint: | |
| """Tests for the /health endpoint.""" | |
| def test_health_returns_200(self): | |
| response = client.get("/health") | |
| assert response.status_code == 200 | |
| def test_health_contains_model_name(self): | |
| response = client.get("/health") | |
| data = response.json() | |
| assert data["model"] == "shiftyspeech" | |
| def test_health_contains_device(self): | |
| response = client.get("/health") | |
| data = response.json() | |
| assert data["device"] == "cpu" | |
| def test_health_contains_status(self): | |
| response = client.get("/health") | |
| data = response.json() | |
| assert data["status"] in ("healthy", "degraded") | |
| class TestPreprocessAudio: | |
| """Tests for audio preprocessing logic.""" | |
| def test_preprocess_short_audio_tiles(self): | |
| """Short audio should be tiled to TARGET_SAMPLES.""" | |
| wav_bytes = _make_wav_bytes(duration_s=0.5, sr=16000) | |
| tensor = api.preprocess_audio(wav_bytes) | |
| assert tensor.shape == (1, api.TARGET_SAMPLES) | |
| def test_preprocess_long_audio_trims(self): | |
| """Long audio should be trimmed to TARGET_SAMPLES.""" | |
| wav_bytes = _make_wav_bytes(duration_s=10.0, sr=16000) | |
| tensor = api.preprocess_audio(wav_bytes) | |
| assert tensor.shape == (1, api.TARGET_SAMPLES) | |
| def test_preprocess_exact_length(self): | |
| """Audio at exact TARGET_SAMPLES should pass through.""" | |
| duration = api.TARGET_SAMPLES / api.SAMPLE_RATE | |
| wav_bytes = _make_wav_bytes(duration_s=duration, sr=16000) | |
| tensor = api.preprocess_audio(wav_bytes) | |
| assert tensor.shape == (1, api.TARGET_SAMPLES) | |
| def test_preprocess_resamples_from_8khz(self): | |
| """Audio at 8kHz should be resampled to 16kHz.""" | |
| wav_bytes = _make_wav_bytes(duration_s=1.0, sr=8000) | |
| tensor = api.preprocess_audio(wav_bytes) | |
| assert tensor.shape == (1, api.TARGET_SAMPLES) | |
| def test_preprocess_invalid_input_raises(self): | |
| """Invalid audio bytes should raise ValueError.""" | |
| with pytest.raises(ValueError): | |
| api.preprocess_audio(b"not audio data") | |
| class TestPredictEndpoint: | |
| """Integration tests for the /predict endpoint (requires weights).""" | |
| def test_predict_returns_200(self): | |
| wav_bytes = _make_wav_bytes(duration_s=2.0) | |
| b64 = base64.b64encode(wav_bytes).decode("utf-8") | |
| response = client.post("/predict", json={"audio_data": b64}) | |
| assert response.status_code == 200 | |
| def test_predict_response_schema(self): | |
| wav_bytes = _make_wav_bytes(duration_s=2.0) | |
| b64 = base64.b64encode(wav_bytes).decode("utf-8") | |
| response = client.post("/predict", json={"audio_data": b64}) | |
| data = response.json() | |
| assert "model" in data | |
| assert "probability" in data | |
| assert "prediction" in data | |
| assert "class" in data | |
| assert "inference_time" in data | |
| assert data["model"] == "shiftyspeech" | |
| def test_predict_probability_in_range(self): | |
| wav_bytes = _make_wav_bytes(duration_s=2.0) | |
| b64 = base64.b64encode(wav_bytes).decode("utf-8") | |
| response = client.post("/predict", json={"audio_data": b64}) | |
| data = response.json() | |
| assert 0.0 <= data["probability"] <= 1.0 | |
| def test_predict_class_matches_prediction(self): | |
| wav_bytes = _make_wav_bytes(duration_s=2.0) | |
| b64 = base64.b64encode(wav_bytes).decode("utf-8") | |
| response = client.post("/predict", json={"audio_data": b64}) | |
| data = response.json() | |
| if data["prediction"] == 1: | |
| assert data["class"] == "fake" | |
| else: | |
| assert data["class"] == "real" | |
| def test_predict_custom_threshold(self): | |
| wav_bytes = _make_wav_bytes(duration_s=2.0) | |
| b64 = base64.b64encode(wav_bytes).decode("utf-8") | |
| response = client.post( | |
| "/predict", | |
| json={"audio_data": b64, "threshold": 0.99}, | |
| ) | |
| data = response.json() | |
| assert response.status_code == 200 | |
| # With threshold=0.99, only very high prob_fake => fake | |
| if data["probability"] < 0.99: | |
| assert data["prediction"] == 0 | |
| assert data["class"] == "real" | |
| def test_predict_inference_time_positive(self): | |
| wav_bytes = _make_wav_bytes(duration_s=2.0) | |
| b64 = base64.b64encode(wav_bytes).decode("utf-8") | |
| response = client.post("/predict", json={"audio_data": b64}) | |
| data = response.json() | |
| assert data["inference_time"] > 0 | |
| class TestRealDataset: | |
| """Integration tests using actual dataset files.""" | |
| def test_predict_real_audio(self): | |
| """Test prediction on a real audio file.""" | |
| path = os.path.join(REAL_DIR, "real_0.wav") | |
| if not os.path.exists(path): | |
| pytest.skip("real_0.wav not found") | |
| b64 = _encode_file(path) | |
| response = client.post("/predict", json={"audio_data": b64}) | |
| assert response.status_code == 200 | |
| data = response.json() | |
| assert 0.0 <= data["probability"] <= 1.0 | |
| def test_predict_fake_audio(self): | |
| """Test prediction on a fake audio file.""" | |
| path = os.path.join(FAKE_DIR, "fake_1.wav") | |
| if not os.path.exists(path): | |
| pytest.skip("fake_1.wav not found") | |
| b64 = _encode_file(path) | |
| response = client.post("/predict", json={"audio_data": b64}) | |
| assert response.status_code == 200 | |
| data = response.json() | |
| assert 0.0 <= data["probability"] <= 1.0 | |
| class TestPredictValidation: | |
| """Tests for input validation on /predict.""" | |
| def test_predict_missing_audio_data(self): | |
| response = client.post("/predict", json={}) | |
| assert response.status_code == 422 | |
| def test_predict_invalid_base64(self): | |
| response = client.post("/predict", json={"audio_data": "not-valid-base64!!!"}) | |
| # Should return 500 (decode error) or 422 | |
| assert response.status_code in (400, 422, 500) | |
| def test_predict_threshold_out_of_range(self): | |
| response = client.post( | |
| "/predict", | |
| json={"audio_data": "dGVzdA==", "threshold": 1.5}, | |
| ) | |
| assert response.status_code == 422 | |