deepsafe's picture
sync from GitHub (0154d02)
4b0b144 verified
Raw History Blame Contribute Delete
10.4 kB
"""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)