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)