"""SafeEar audio deepfake detection API service. Uses the SafeEar content privacy-preserving model (CCS 2024) to detect synthetic speech. Two-stage pipeline: 1. SpeechTokenizer (neural audio codec) decouples acoustic features 2. SafeEar1s (transformer classifier) detects spoofing from acoustic tokens Weights: HuggingFace TEC2004/SafeEar-ASV19-spoof-detection """ import base64 import logging import os import sys import tempfile import time from typing import Optional import uvicorn from fastapi import FastAPI, HTTPException from pydantic import BaseModel, Field logging.basicConfig( level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", ) logger = logging.getLogger("safeear_api") import platform import librosa import numpy as np import torch # Add SafeEar repo to path for model imports SAFEEAR_REPO_PATH = os.environ.get( "SAFEEAR_REPO_PATH", os.path.join(os.path.dirname(__file__), "safeear_repo"), ) if SAFEEAR_REPO_PATH not in sys.path: sys.path.insert(0, SAFEEAR_REPO_PATH) 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 ( 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") # Constants MODEL_NAME = "safeear" WEIGHTS_DIR = os.environ.get( "WEIGHTS_DIR", os.path.join(os.path.dirname(__file__), "weights"), ) 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 MAX_AUDIO_LENGTH = 64600 # ~4 seconds at 16kHz (ASVspoof standard) SOFTMAX_TEMPERATURE = 5.0 # Calibration temperature for out-of-distribution data NUM_INFERENCE_PASSES = 5 # Monte Carlo passes for stable predictions # Global model instances decouple_model = None detect_model = None class AudioInput(BaseModel): """Schema for audio prediction requests.""" 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="SafeEar Audio Deepfake Detection API", description="Content privacy-preserving deepfake detection using SafeEar.", version="1.0.0", ) def load_models(): """Load both the decouple model (SpeechTokenizer) and detect model.""" global decouple_model, detect_model if decouple_model is not None and detect_model is not None: return True speech_tokenizer_path = os.path.join(WEIGHTS_DIR, "SpeechTokenizer.pt") checkpoint_path = os.path.join(WEIGHTS_DIR, "model.ckpt") if not os.path.exists(speech_tokenizer_path): logger.error(f"SpeechTokenizer weights not found: {speech_tokenizer_path}") return False if not os.path.exists(checkpoint_path): logger.error(f"Model checkpoint not found: {checkpoint_path}") return False try: # --- Load SpeechTokenizer (decouple model) --- from safeear.models.decouple import SpeechTokenizer logger.info("Loading SpeechTokenizer...") decouple_model = SpeechTokenizer( n_filters=64, strides=[8, 5, 4, 2], dimension=1024, semantic_dimension=768, bidirectional=True, dilation_base=2, residual_kernel_size=3, n_residual_layers=1, lstm_layers=2, activation="ELU", codebook_size=1024, n_q=8, sample_rate=16000, ) st_state = torch.load(speech_tokenizer_path, map_location="cpu") decouple_model.load_state_dict(st_state) decouple_model.to(DEVICE) decouple_model.eval() logger.info("SpeechTokenizer loaded.") # --- Load SafeEar1s (detect model) from Lightning checkpoint --- from safeear.models.safeear import SafeEar1s, SE_Rawformer_front logger.info("Loading SafeEar1s detect model...") detect_model = SafeEar1s( front=SE_Rawformer_front(), embedding_dim=1024, dropout_rate=0.1, attention_dropout=0.1, stochastic_depth=0.1, num_layers=2, num_heads=8, num_classes=2, positional_embedding="sine", mlp_ratio=1.0, ) # The .ckpt is a PyTorch Lightning checkpoint ckpt = torch.load(checkpoint_path, map_location="cpu") state_dict = ckpt.get("state_dict", ckpt) # Lightning prefixes keys with "detect_model." detect_state = {} for k, v in state_dict.items(): if k.startswith("detect_model."): detect_state[k.replace("detect_model.", "", 1)] = v detect_model.load_state_dict(detect_state) detect_model.to(DEVICE) detect_model.eval() logger.info("SafeEar1s detect model loaded.") return True except Exception as e: logger.exception(f"Failed to load SafeEar models: {e}") decouple_model = None detect_model = None return False def preprocess_audio(audio_bytes: bytes) -> torch.Tensor: """Load audio bytes, resample to 16kHz mono, pad/trim.""" with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: tmp.write(audio_bytes) tmp_path = tmp.name try: waveform, _ = librosa.load(tmp_path, sr=SAMPLE_RATE, mono=True) finally: os.unlink(tmp_path) if len(waveform) < MAX_AUDIO_LENGTH: waveform = np.pad(waveform, (0, MAX_AUDIO_LENGTH - len(waveform))) else: waveform = waveform[:MAX_AUDIO_LENGTH] # Shape: (1, 1, samples) -- batch=1, channels=1, time tensor = torch.FloatTensor(waveform).unsqueeze(0).unsqueeze(0).to(DEVICE) return tensor @app.on_event("startup") async def startup_event(): """Attempt to load models at startup.""" load_models() 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 {} @app.get("/health") async def health(): """Return service health status and model availability.""" models_loaded = decouple_model is not None and detect_model is not None return { "status": "healthy" if models_loaded else "degraded", "model": MODEL_NAME, "device": str(DEVICE), "weights_found": ( os.path.exists(os.path.join(WEIGHTS_DIR, "SpeechTokenizer.pt")) and os.path.exists(os.path.join(WEIGHTS_DIR, "model.ckpt")) ), **_gpu_health_info(), } @app.post("/predict") async def predict(input_data: AudioInput): """Run SafeEar inference on base64-encoded audio data.""" if decouple_model is None or detect_model is None: if not load_models(): raise HTTPException(status_code=503, detail="Models not loaded") try: start_time = time.time() logger.info( "Received prediction request. " f"Data size: {len(input_data.audio_data)} chars" ) audio_bytes = base64.b64decode(input_data.audio_data) x_wav = preprocess_audio(audio_bytes) with torch.no_grad(): # Step 1: Extract acoustic tokens via SpeechTokenizer # forward() returns: # (reconstructed, commit_loss, semantic_feature, acoustic_tokens) # layers=[0,1,2,3,4,5,6,7] means layer 0 goes to # semantic_feature; layers 1-7 go to acoustic_tokens list _, _, _, acoustic_tokens = decouple_model( x_wav, layers=[0, 1, 2, 3, 4, 5, 6, 7] ) # Step 2: Run detection model with Monte Carlo averaging # SafeEar1s uses torch.randperm() in forward, so we average # multiple passes for stable predictions logit_sum = torch.zeros(1, 2, device=DEVICE) for _ in range(NUM_INFERENCE_PASSES): raw_logits, _ = detect_model(acoustic_tokens) logit_sum += raw_logits avg_logits = logit_sum / NUM_INFERENCE_PASSES # Step 3: Get fake probability with temperature-scaled softmax # The model produces extreme logits that saturate standard # softmax. Temperature scaling preserves discrimination while # giving more interpretable probabilities. probs = torch.softmax(avg_logits / SOFTMAX_TEMPERATURE, dim=-1) prob_fake = probs[0, 1].item() prediction = 1 if prob_fake >= input_data.threshold else 0 verdict = "fake" if prediction == 1 else "real" inference_time = time.time() - start_time 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", 8002)) uvicorn.run(app, host="0.0.0.0", port=port)