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/api.py from deepsafe/deepsafe-services: direct link, hf CLI and curl.
- Browser
- Download file 9.31 kB
-
https://huggingface.co/deepsafe/deepsafe-services/resolve/main/audio/shiftyspeech/api.py
- Command line
-
hf download hf://deepsafe/deepsafe-services/audio/shiftyspeech/api.py
-
curl -L -o api.py https://huggingface.co/deepsafe/deepsafe-services/resolve/main/audio/shiftyspeech/api.py
9.31 kB
| """ShiftySpeech (SSL-AASIST) Audio Deepfake Detection API. | |
| Detects synthetic speech using the SSL-AASIST model architecture: | |
| - Frontend: XLSR wav2vec 2.0 (Self-Supervised Learning) | |
| - Backend: AASIST (Audio Anti-Spoofing using Integrated | |
| Spectro-Temporal Graph Attention Networks) | |
| Reference: https://github.com/Ashigarg123/ShiftySpeech | |
| """ | |
| import base64 | |
| import io | |
| import logging | |
| import os | |
| import platform | |
| import sys | |
| import time | |
| import warnings | |
| from typing import Optional | |
| import librosa | |
| import numpy as np | |
| import torch | |
| import uvicorn | |
| from fastapi import FastAPI, HTTPException | |
| from pydantic import BaseModel, Field | |
| # Suppress deprecation warnings from fairseq/omegaconf compatibility | |
| warnings.filterwarnings("ignore", category=DeprecationWarning) | |
| # Monkey-patch omegaconf for fairseq compatibility (older fairseq | |
| # expects is_primitive_type which was removed in newer omegaconf). | |
| 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) | |
| # Configure logging | |
| logging.basicConfig( | |
| level=logging.INFO, | |
| format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", | |
| ) | |
| logger = logging.getLogger("shiftyspeech_api") | |
| # Add the SSL_Anti-spoofing model code to the path | |
| MODEL_CODE_PATH = "/app/synthetic_speech_detection/SSL_Anti-spoofing" | |
| if MODEL_CODE_PATH not in sys.path: | |
| sys.path.insert(0, MODEL_CODE_PATH) | |
| # Import model class (deferred to allow path setup) | |
| try: | |
| from model import Model as SSLAASISTModel | |
| except ImportError as e: | |
| logger.error(f"Failed to import SSL-AASIST model: {e}") | |
| SSLAASISTModel = None | |
| # Constants | |
| MODEL_NAME = "shiftyspeech" | |
| MODEL_ID = "ssl_aasist_augmented" | |
| WEIGHTS_PATH = "/app/weights/hfg_aug_1_2.pt" | |
| XLSR_DIR = "/app/models" | |
| def _get_device(): | |
| """Select optimal device: MPS (Apple) > CUDA (NVIDIA) > CPU.""" | |
| override = os.environ.get("DEEPSAFE_DEVICE", "").strip().lower() | |
| if override == "cpu": | |
| return torch.device("cpu") | |
| if override == "cuda" and torch.cuda.is_available(): | |
| return torch.device("cuda") | |
| if ( | |
| override == "mps" | |
| and hasattr(torch.backends, "mps") | |
| and torch.backends.mps.is_available() | |
| ): | |
| return torch.device("mps") | |
| if override: | |
| pass # Invalid override, fall through to auto-detect | |
| if ( | |
| platform.system() == "Darwin" | |
| and hasattr(torch.backends, "mps") | |
| and torch.backends.mps.is_available() | |
| ): | |
| return torch.device("mps") | |
| if torch.cuda.is_available(): | |
| return torch.device("cuda") | |
| return torch.device("cpu") | |
| DEVICE = _get_device() | |
| if DEVICE.type == "cuda": | |
| torch.backends.cudnn.benchmark = True | |
| torch.set_float32_matmul_precision("high") | |
| if DEVICE.type == "cuda": | |
| logger.info( | |
| "Device: cuda (%s, %.1f GB VRAM)", | |
| torch.cuda.get_device_name(0), | |
| torch.cuda.get_device_properties(0).total_memory / 1024**3, | |
| ) | |
| else: | |
| logger.warning( | |
| "Device: %s (no CUDA available -- check nvidia-container-toolkit)", | |
| DEVICE, | |
| ) | |
| SAMPLE_RATE = 16000 | |
| TARGET_SAMPLES = 64600 # ~4.04 seconds at 16kHz | |
| # Global model instance | |
| model = None | |
| class AudioInput(BaseModel): | |
| """Request schema for audio deepfake detection.""" | |
| audio_data: str = Field( | |
| ..., description="Base64 encoded audio string (WAV/MP3/etc)" | |
| ) | |
| threshold: Optional[float] = Field( | |
| 0.5, ge=0.0, le=1.0, description="Classification threshold" | |
| ) | |
| app = FastAPI( | |
| title="ShiftySpeech Audio Deepfake Detection API", | |
| description=( | |
| "Service for detecting synthetic speech using the " | |
| "SSL-AASIST model (XLSR wav2vec 2.0 + AASIST backend)." | |
| ), | |
| version="1.0.0", | |
| ) | |
| def load_model(): | |
| """Load the SSL-AASIST model with augmented weights. | |
| Returns: | |
| The loaded model, or None if loading fails. | |
| """ | |
| global model | |
| if model is not None: | |
| return model | |
| logger.info(f"Loading SSL-AASIST model onto {DEVICE}...") | |
| if SSLAASISTModel is None: | |
| logger.error("SSL-AASIST model class not available.") | |
| return None | |
| if not os.path.exists(WEIGHTS_PATH): | |
| logger.error(f"Model weights not found at {WEIGHTS_PATH}") | |
| return None | |
| try: | |
| # Ensure XLSR model directory exists for architecture init | |
| os.makedirs(XLSR_DIR, exist_ok=True) | |
| import argparse | |
| args = argparse.Namespace() | |
| model = SSLAASISTModel(args, str(DEVICE)) | |
| # Load fine-tuned weights (includes XLSR weights) | |
| try: | |
| state_dict = torch.load( | |
| WEIGHTS_PATH, | |
| map_location=DEVICE, | |
| weights_only=False, | |
| ) | |
| except TypeError: | |
| state_dict = torch.load(WEIGHTS_PATH, map_location=DEVICE) | |
| model.load_state_dict(state_dict) | |
| model.to(DEVICE) | |
| model.eval() | |
| logger.info("SSL-AASIST model loaded successfully.") | |
| return model | |
| except Exception as e: | |
| logger.exception(f"Failed to load SSL-AASIST model: {e}") | |
| model = None | |
| return None | |
| async def startup_event(): | |
| """Load model on service startup.""" | |
| load_model() | |
| def _gpu_health_info() -> dict: | |
| """Return GPU metrics for the health endpoint.""" | |
| if torch.cuda.is_available() and DEVICE.type == "cuda": | |
| return { | |
| "gpu_name": torch.cuda.get_device_name(0), | |
| "vram_used_mb": round(torch.cuda.memory_allocated(0) / 1024**2), | |
| "vram_total_mb": round( | |
| torch.cuda.get_device_properties(0).total_memory / 1024**2 | |
| ), | |
| } | |
| return {} | |
| async def health(): | |
| """Health check endpoint.""" | |
| return { | |
| "status": "healthy" if model is not None else "degraded", | |
| "model": MODEL_NAME, | |
| "model_id": MODEL_ID, | |
| "device": str(DEVICE), | |
| "weights_found": os.path.exists(WEIGHTS_PATH), | |
| **_gpu_health_info(), | |
| } | |
| def preprocess_audio(audio_bytes: bytes) -> torch.Tensor: | |
| """Preprocess audio for SSL-AASIST inference. | |
| Loads audio, resamples to 16kHz mono, and pads/trims | |
| to TARGET_SAMPLES using tiling (matching original training | |
| preprocessing from data_utils.py). | |
| Args: | |
| audio_bytes: Raw audio file bytes. | |
| Returns: | |
| Audio tensor of shape (1, TARGET_SAMPLES). | |
| Raises: | |
| ValueError: If audio preprocessing fails. | |
| """ | |
| try: | |
| logger.info("Starting audio preprocessing...") | |
| audio, sr = librosa.load(io.BytesIO(audio_bytes), sr=SAMPLE_RATE, mono=True) | |
| logger.info(f"Audio loaded. Length: {len(audio)} samples at {sr}Hz") | |
| # Pad/trim to TARGET_SAMPLES using tiling | |
| # (matches original data_utils.pad function) | |
| if len(audio) >= TARGET_SAMPLES: | |
| audio = audio[:TARGET_SAMPLES] | |
| else: | |
| num_repeats = TARGET_SAMPLES // len(audio) + 1 | |
| audio = np.tile(audio, num_repeats)[:TARGET_SAMPLES] | |
| logger.info(f"Audio padded/trimmed to {TARGET_SAMPLES} samples") | |
| audio_tensor = torch.FloatTensor(audio).unsqueeze(0).to(DEVICE) | |
| return audio_tensor | |
| except Exception as e: | |
| logger.error(f"Error preprocessing audio: {e}") | |
| raise ValueError(f"Audio preprocessing failed: {str(e)}") | |
| async def predict(input_data: AudioInput): | |
| """Run deepfake detection on base64-encoded audio. | |
| The model outputs 2 logits: [spoof_score, bonafide_score]. | |
| Class 0 = spoof (fake), Class 1 = bonafide (real). | |
| The returned probability is the spoof/fake probability. | |
| """ | |
| if model is None: | |
| if load_model() is None: | |
| raise HTTPException(status_code=503, detail="Model not loaded") | |
| try: | |
| start_time = time.time() | |
| logger.info( | |
| f"Prediction request. Data size: " f"{len(input_data.audio_data)} chars" | |
| ) | |
| # Decode base64 audio | |
| audio_bytes = base64.b64decode(input_data.audio_data) | |
| # Preprocess | |
| audio_tensor = preprocess_audio(audio_bytes) | |
| # Inference | |
| logger.info("Starting model inference...") | |
| with torch.no_grad(): | |
| output = model(audio_tensor) | |
| # output shape: [batch, 2] | |
| # Index 0 = spoof logit, Index 1 = bonafide logit | |
| probs = torch.softmax(output, dim=1) | |
| prob_fake = probs[0, 0].item() | |
| prediction = 1 if prob_fake >= input_data.threshold else 0 | |
| verdict = "fake" if prediction == 1 else "real" | |
| inference_time = time.time() - start_time | |
| logger.info( | |
| f"Prediction: {verdict} (prob_fake={prob_fake:.4f}, " | |
| f"time={inference_time:.3f}s)" | |
| ) | |
| return { | |
| "model": MODEL_NAME, | |
| "probability": float(prob_fake), | |
| "prediction": int(prediction), | |
| "class": verdict, | |
| "inference_time": float(inference_time), | |
| } | |
| except Exception as e: | |
| logger.exception(f"Error during prediction: {e}") | |
| raise HTTPException(status_code=500, detail=str(e)) | |
| if __name__ == "__main__": | |
| port = int(os.environ.get("MODEL_PORT", 8001)) | |
| uvicorn.run(app, host="0.0.0.0", port=port) | |