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
File size: 10,413 Bytes
4b0b144 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 | """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)
|