deepsafe commited on
Commit
4b0b144
Β·
verified Β·
1 Parent(s): a29616b

sync from GitHub (0154d02)

Browse files
This view is limited to 50 files because it contains too many changes. Β  See raw diff
Files changed (50) hide show
  1. audio/aasist3/.gitignore +1 -0
  2. audio/aasist3/Dockerfile +35 -0
  3. audio/aasist3/api.py +345 -0
  4. audio/aasist3/model/__init__.py +1 -0
  5. audio/aasist3/model/branch.py +34 -0
  6. audio/aasist3/model/full_model.py +139 -0
  7. audio/aasist3/model/gat.py +99 -0
  8. audio/aasist3/model/hs_gal.py +176 -0
  9. audio/aasist3/model/kan.py +213 -0
  10. audio/aasist3/model/pool.py +45 -0
  11. audio/aasist3/model/residual.py +56 -0
  12. audio/aasist3/model/wav2vec.py +82 -0
  13. audio/aasist3/requirements.txt +11 -0
  14. audio/nes2net/.gitignore +1 -0
  15. audio/nes2net/Dockerfile +41 -0
  16. audio/nes2net/api.py +315 -0
  17. audio/nes2net/model_scripts/__init__.py +0 -0
  18. audio/nes2net/model_scripts/wav2vec2_Nes2Net_X.py +317 -0
  19. audio/nes2net/requirements.txt +13 -0
  20. audio/safeear/Dockerfile +49 -0
  21. audio/safeear/api.py +321 -0
  22. audio/safeear/download_weights.sh +23 -0
  23. audio/safeear/requirements.txt +15 -0
  24. audio/shiftyspeech/Dockerfile +43 -0
  25. audio/shiftyspeech/api.py +315 -0
  26. audio/shiftyspeech/evaluate.py +210 -0
  27. audio/shiftyspeech/requirements.txt +14 -0
  28. audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/.env +2 -0
  29. audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/LICENSE +21 -0
  30. audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/RawBoost.py +143 -0
  31. audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/Simplified_CM_solution.py +227 -0
  32. audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/data_utils.py +292 -0
  33. audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/model.py +603 -0
  34. audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/requirements.txt +4 -0
  35. audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/startup_config.py +60 -0
  36. audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/train.py +446 -0
  37. audio/shiftyspeech/tests/__init__.py +0 -0
  38. audio/shiftyspeech/tests/test_api.py +253 -0
  39. audio/sonics/Dockerfile +36 -0
  40. audio/sonics/app.py +285 -0
  41. audio/sonics/requirements.txt +22 -0
  42. ensemble-core/Dockerfile +10 -0
  43. ensemble-core/main.py +23 -0
  44. ensemble-core/requirements.txt +1 -0
  45. ensemble-core/scripts/create_dataset.py +169 -0
  46. ensemble-core/scripts/meta_feature_generator.py +366 -0
  47. ensemble-core/scripts/train_meta_learner_advanced.py +1228 -0
  48. image/aide/.gitignore +5 -0
  49. image/aide/Dockerfile +47 -0
  50. image/aide/app.py +433 -0
audio/aasist3/.gitignore ADDED
@@ -0,0 +1 @@
 
 
1
+ weights/
audio/aasist3/Dockerfile ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM nvidia/cuda:12.1.1-cudnn8-runtime-ubuntu22.04
2
+
3
+ ENV DEBIAN_FRONTEND=noninteractive
4
+ ENV PYTHONUNBUFFERED=1
5
+
6
+ RUN apt-get update && apt-get install -y --no-install-recommends \
7
+ python3 python3-pip \
8
+ ffmpeg libsndfile1 \
9
+ && rm -rf /var/lib/apt/lists/*
10
+
11
+ RUN ln -sf /usr/bin/python3 /usr/bin/python
12
+
13
+ WORKDIR /app
14
+
15
+ # Install PyTorch with CUDA 12.1
16
+ RUN pip install --no-cache-dir \
17
+ torch==2.5.1 torchaudio==2.5.1 \
18
+ --index-url https://download.pytorch.org/whl/cu121
19
+
20
+ COPY requirements.txt .
21
+ RUN pip install --no-cache-dir -r requirements.txt
22
+
23
+ RUN python -c "from transformers import Wav2Vec2Config; Wav2Vec2Config.from_pretrained('facebook/wav2vec2-large-xlsr-53', cache_dir='/app/w2v_cache')"
24
+
25
+ COPY model/ /app/model/
26
+ COPY api.py .
27
+ RUN mkdir -p /app/weights
28
+ COPY weights/ /app/weights/
29
+
30
+ EXPOSE 8005
31
+
32
+ RUN adduser --disabled-password --gecos '' appuser
33
+ USER appuser
34
+
35
+ CMD ["python", "api.py"]
audio/aasist3/api.py ADDED
@@ -0,0 +1,345 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """AASIST3 Audio Deepfake Detection API.
2
+
3
+ Detects synthetic speech using the AASIST3 model architecture:
4
+ - Frontend: XLSR wav2vec 2.0 (HuggingFace Transformers)
5
+ - Backend: AASIST with KAN (Kolmogorov-Arnold Network) linear
6
+ layers and Graph Attention Networks
7
+
8
+ Reference: https://github.com/AI4Bharat/AASIST3
9
+ """
10
+
11
+ import base64
12
+ import io
13
+ import logging
14
+ import os
15
+ import platform
16
+ import time
17
+ from typing import Optional
18
+
19
+ import numpy as np
20
+ import soundfile as sf
21
+ import torch
22
+ import torch.nn.functional as F
23
+ import uvicorn
24
+ from fastapi import FastAPI, HTTPException
25
+ from pydantic import BaseModel, Field
26
+
27
+ # Point transformers cache to pre-cached wav2vec2 config
28
+ # (must be set before importing model code)
29
+ os.environ["TRANSFORMERS_CACHE"] = "/app/w2v_cache"
30
+
31
+ # Configure logging
32
+ logging.basicConfig(
33
+ level=logging.INFO,
34
+ format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
35
+ )
36
+ logger = logging.getLogger("aasist3_api")
37
+
38
+ # Import model class
39
+ try:
40
+ from model import aasist3 as AASIST3Model
41
+ except ImportError as e:
42
+ logger.error(f"Failed to import AASIST3 model: {e}")
43
+ AASIST3Model = None
44
+
45
+ # Constants
46
+ MODEL_NAME = "aasist3"
47
+ MODEL_ID = "aasist3_kan_mlaad"
48
+ WEIGHTS_DIR = "/app/weights"
49
+
50
+
51
+ def _get_device():
52
+ """Select optimal device: MPS (Apple) > CUDA (NVIDIA) > CPU."""
53
+ override = os.environ.get("DEEPSAFE_DEVICE", "").strip().lower()
54
+ if override == "cpu":
55
+ return torch.device("cpu")
56
+ if override == "cuda" and torch.cuda.is_available():
57
+ return torch.device("cuda")
58
+ if (
59
+ override == "mps"
60
+ and hasattr(torch.backends, "mps")
61
+ and torch.backends.mps.is_available()
62
+ ):
63
+ return torch.device("mps")
64
+ if override:
65
+ pass # Invalid override, fall through to auto-detect
66
+ if (
67
+ platform.system() == "Darwin"
68
+ and hasattr(torch.backends, "mps")
69
+ and torch.backends.mps.is_available()
70
+ ):
71
+ return torch.device("mps")
72
+ if torch.cuda.is_available():
73
+ return torch.device("cuda")
74
+ return torch.device("cpu")
75
+
76
+
77
+ DEVICE = _get_device()
78
+
79
+ if DEVICE.type == "cuda":
80
+ torch.backends.cudnn.benchmark = True
81
+ torch.set_float32_matmul_precision("high")
82
+
83
+ if DEVICE.type == "cuda":
84
+ logger.info(
85
+ "Device: cuda (%s, %.1f GB VRAM)",
86
+ torch.cuda.get_device_name(0),
87
+ torch.cuda.get_device_properties(0).total_memory / 1024**3,
88
+ )
89
+ else:
90
+ logger.warning(
91
+ "Device: %s (no CUDA available -- check nvidia-container-toolkit)",
92
+ DEVICE,
93
+ )
94
+
95
+ SAMPLE_RATE = 16000
96
+ TARGET_SAMPLES = 64600 # ~4.04 seconds at 16kHz
97
+
98
+ # Global model instance
99
+ model = None
100
+
101
+
102
+ class AudioInput(BaseModel):
103
+ """Request schema for audio deepfake detection."""
104
+
105
+ audio_data: str = Field(
106
+ ..., description="Base64 encoded audio string (WAV/MP3/etc)"
107
+ )
108
+ threshold: Optional[float] = Field(
109
+ 0.5, ge=0.0, le=1.0, description="Classification threshold"
110
+ )
111
+
112
+
113
+ app = FastAPI(
114
+ title="AASIST3 Audio Deepfake Detection API",
115
+ description=(
116
+ "Service for detecting synthetic speech using the "
117
+ "AASIST3 model (HuggingFace wav2vec 2.0 + AASIST "
118
+ "with KAN layers)."
119
+ ),
120
+ version="1.0.0",
121
+ )
122
+
123
+
124
+ def load_model():
125
+ """Load the AASIST3 model from pretrained weights.
126
+
127
+ Returns:
128
+ The loaded model, or None if loading fails.
129
+ """
130
+ global model
131
+ if model is not None:
132
+ return model
133
+
134
+ logger.info(f"Loading AASIST3 model onto {DEVICE}...")
135
+
136
+ if AASIST3Model is None:
137
+ logger.error("AASIST3 model class not available.")
138
+ return None
139
+
140
+ weights_safetensors = os.path.join(WEIGHTS_DIR, "model.safetensors")
141
+ weights_config = os.path.join(WEIGHTS_DIR, "config.json")
142
+
143
+ if not os.path.exists(weights_safetensors):
144
+ logger.error(f"Model weights not found at {weights_safetensors}")
145
+ return None
146
+
147
+ if not os.path.exists(weights_config):
148
+ logger.error(f"Model config not found at {weights_config}")
149
+ return None
150
+
151
+ try:
152
+ model = AASIST3Model.from_pretrained(WEIGHTS_DIR)
153
+ model.to(DEVICE)
154
+ # Set model to inference mode (disables dropout, batchnorm)
155
+ model.train(False)
156
+
157
+ logger.info("AASIST3 model loaded successfully.")
158
+ return model
159
+ except Exception as e:
160
+ logger.exception(f"Failed to load AASIST3 model: {e}")
161
+ model = None
162
+ return None
163
+
164
+
165
+ @app.on_event("startup")
166
+ async def startup_event():
167
+ """Load model on service startup."""
168
+ load_model()
169
+
170
+
171
+ def _gpu_health_info() -> dict:
172
+ """Return GPU metrics for the health endpoint."""
173
+ if torch.cuda.is_available() and DEVICE.type == "cuda":
174
+ return {
175
+ "gpu_name": torch.cuda.get_device_name(0),
176
+ "vram_used_mb": round(torch.cuda.memory_allocated(0) / 1024**2),
177
+ "vram_total_mb": round(
178
+ torch.cuda.get_device_properties(0).total_memory / 1024**2
179
+ ),
180
+ }
181
+ return {}
182
+
183
+
184
+ @app.get("/health")
185
+ async def health():
186
+ """Health check endpoint."""
187
+ return {
188
+ "status": "healthy" if model is not None else "degraded",
189
+ "model": MODEL_NAME,
190
+ "model_id": MODEL_ID,
191
+ "device": str(DEVICE),
192
+ "weights_found": os.path.exists(os.path.join(WEIGHTS_DIR, "model.safetensors")),
193
+ **_gpu_health_info(),
194
+ }
195
+
196
+
197
+ def _load_audio_bytes(audio_bytes: bytes) -> tuple:
198
+ """Load audio from raw bytes using soundfile with torchaudio fallback.
199
+
200
+ Args:
201
+ audio_bytes: Raw audio file bytes.
202
+
203
+ Returns:
204
+ Tuple of (audio_numpy_array, sample_rate).
205
+
206
+ Raises:
207
+ ValueError: If audio cannot be loaded by any backend.
208
+ """
209
+ # Try soundfile first (handles WAV, FLAC natively)
210
+ sf_error = None
211
+ try:
212
+ audio, sr = sf.read(io.BytesIO(audio_bytes), dtype="float32")
213
+ if audio.ndim > 1:
214
+ audio = audio.mean(axis=1) # Convert to mono
215
+ return audio, sr
216
+ except Exception as sf_err:
217
+ sf_error = sf_err
218
+ logger.debug(f"soundfile failed, trying torchaudio: {sf_err}")
219
+
220
+ # Fallback to torchaudio (handles MP3, compressed formats)
221
+ try:
222
+ import torchaudio
223
+
224
+ buf = io.BytesIO(audio_bytes)
225
+ waveform, sr = torchaudio.load(buf)
226
+ if waveform.shape[0] > 1:
227
+ waveform = waveform.mean(dim=0, keepdim=True)
228
+ return waveform.squeeze(0).numpy(), sr
229
+ except Exception as ta_err:
230
+ raise ValueError(
231
+ f"Failed to load audio with soundfile and torchaudio: "
232
+ f"sf={sf_error}, ta={ta_err}"
233
+ )
234
+
235
+
236
+ def preprocess_audio(audio_bytes: bytes) -> torch.Tensor:
237
+ """Preprocess audio for AASIST3 inference.
238
+
239
+ Loads audio, resamples to 16kHz mono, and zero-pads or
240
+ truncates to TARGET_SAMPLES.
241
+
242
+ Args:
243
+ audio_bytes: Raw audio file bytes.
244
+
245
+ Returns:
246
+ Audio tensor of shape (1, TARGET_SAMPLES).
247
+
248
+ Raises:
249
+ ValueError: If audio preprocessing fails.
250
+ """
251
+ try:
252
+ logger.info("Starting audio preprocessing...")
253
+ audio, sr = _load_audio_bytes(audio_bytes)
254
+ logger.info(f"Audio loaded. Length: {len(audio)} samples at {sr}Hz")
255
+
256
+ # Resample to 16kHz if needed
257
+ if sr != SAMPLE_RATE:
258
+ import torchaudio
259
+
260
+ resampler = torchaudio.transforms.Resample(
261
+ orig_freq=sr, new_freq=SAMPLE_RATE
262
+ )
263
+ audio_tensor = torch.FloatTensor(audio).unsqueeze(0)
264
+ audio_tensor = resampler(audio_tensor).squeeze(0)
265
+ audio = audio_tensor.numpy()
266
+ logger.info(
267
+ f"Resampled from {sr}Hz to {SAMPLE_RATE}Hz. "
268
+ f"New length: {len(audio)} samples"
269
+ )
270
+
271
+ # Zero-pad or truncate to TARGET_SAMPLES
272
+ if len(audio) >= TARGET_SAMPLES:
273
+ audio = audio[:TARGET_SAMPLES]
274
+ else:
275
+ pad_length = TARGET_SAMPLES - len(audio)
276
+ audio = np.pad(audio, (0, pad_length), mode="constant")
277
+
278
+ logger.info(f"Audio padded/trimmed to {TARGET_SAMPLES} samples")
279
+
280
+ audio_tensor = torch.FloatTensor(audio).unsqueeze(0).to(DEVICE)
281
+ return audio_tensor
282
+ except Exception as e:
283
+ logger.error(f"Error preprocessing audio: {e}")
284
+ raise ValueError(f"Audio preprocessing failed: {str(e)}")
285
+
286
+
287
+ @app.post("/predict")
288
+ async def predict(input_data: AudioInput):
289
+ """Run deepfake detection on base64-encoded audio.
290
+
291
+ The model outputs 2 logits: [bonafide_score, spoof_score].
292
+ Class 0 = bonafide (real), Class 1 = spoof (fake).
293
+ The returned probability is the spoof/fake probability.
294
+ """
295
+ if model is None:
296
+ if load_model() is None:
297
+ raise HTTPException(status_code=503, detail="Model not loaded")
298
+
299
+ try:
300
+ start_time = time.time()
301
+ logger.info(
302
+ f"Prediction request. Data size: " f"{len(input_data.audio_data)} chars"
303
+ )
304
+
305
+ # Decode base64 audio
306
+ audio_bytes = base64.b64decode(input_data.audio_data)
307
+
308
+ # Preprocess
309
+ audio_tensor = preprocess_audio(audio_bytes)
310
+
311
+ # Inference
312
+ logger.info("Starting model inference...")
313
+ with torch.no_grad():
314
+ output = model(audio_tensor)
315
+
316
+ # output shape: [batch, 2]
317
+ # Index 0 = bonafide logit, Index 1 = spoof logit
318
+ probs = torch.softmax(output, dim=1)
319
+ prob_fake = probs[0, 1].item()
320
+
321
+ prediction = 1 if prob_fake >= input_data.threshold else 0
322
+ verdict = "fake" if prediction == 1 else "real"
323
+ inference_time = time.time() - start_time
324
+
325
+ logger.info(
326
+ f"Prediction: {verdict} (prob_fake={prob_fake:.4f}, "
327
+ f"time={inference_time:.3f}s)"
328
+ )
329
+
330
+ return {
331
+ "model": MODEL_NAME,
332
+ "probability": float(prob_fake),
333
+ "prediction": int(prediction),
334
+ "class": verdict,
335
+ "inference_time": float(inference_time),
336
+ }
337
+
338
+ except Exception as e:
339
+ logger.exception(f"Error during prediction: {e}")
340
+ raise HTTPException(status_code=500, detail=str(e))
341
+
342
+
343
+ if __name__ == "__main__":
344
+ port = int(os.environ.get("MODEL_PORT", 8005))
345
+ uvicorn.run(app, host="0.0.0.0", port=port)
audio/aasist3/model/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .full_model import aasist3
audio/aasist3/model/branch.py ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch.nn as nn
2
+
3
+ from .hs_gal import HtrgGraphAttentionLayer
4
+ from .pool import GraphPool
5
+
6
+ class InferenceBranch(nn.Module):
7
+ def __init__(self, gat_dims, temperature, pool_ratio, size):
8
+ super().__init__()
9
+ self.htrg_gat1 = HtrgGraphAttentionLayer(
10
+ gat_dims[0], gat_dims[1], temperature=temperature, size=size
11
+ )
12
+ self.htrg_gat2 = HtrgGraphAttentionLayer(
13
+ gat_dims[1], gat_dims[1], temperature=temperature, size=size
14
+ )
15
+
16
+ self.pool_hS = GraphPool(pool_ratio, gat_dims[1], 0.3, size=size)
17
+ self.pool_hT = GraphPool(pool_ratio, gat_dims[1], 0.3, size=size)
18
+
19
+ def forward(self, out_T, out_S, master):
20
+ # ΠŸΠ΅Ρ€Π²Π°Ρ стадия
21
+ out_T_res, out_S_res, master_res = self.htrg_gat1(out_T, out_S, master=master)
22
+
23
+ # ΠŸΡƒΠ»ΠΈΠ½Π³
24
+ out_S_res = self.pool_hS(out_S_res)
25
+ out_T_res = self.pool_hT(out_T_res)
26
+
27
+ # Вторая стадия с residual connection
28
+ out_T_aug, out_S_aug, master_aug = self.htrg_gat2(out_T_res, out_S_res, master=master_res)
29
+
30
+ out_T_final = out_T_res + out_T_aug
31
+ out_S_final = out_S_res + out_S_aug
32
+ master_final = master_res + master_aug
33
+
34
+ return out_T_final, out_S_final, master_final
audio/aasist3/model/full_model.py ADDED
@@ -0,0 +1,139 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+
3
+ import torch
4
+ import torch.nn.functional as F
5
+ import torch.nn as nn
6
+ from huggingface_hub import PyTorchModelHubMixin
7
+
8
+ from .kan import KANLinear
9
+ from .gat import GraphAttentionLayer
10
+ from .pool import GraphPool
11
+ from .branch import InferenceBranch
12
+ from .residual import Residual_block
13
+ from .wav2vec import Wav2Vec2Encoder
14
+
15
+
16
+ class aasist3(nn.Module, PyTorchModelHubMixin):
17
+ def __init__(self, d_args={
18
+ "architecture": "AASIST",
19
+ "nb_samp": 64600,
20
+ "first_conv": 128,
21
+ "filts": [70, [1, 32], [32, 32], [32, 64], [64, 64]],
22
+ "gat_dims": [64, 32],
23
+ "pool_ratios": [0.5, 0.7, 0.5, 0.5],
24
+ "temperatures": [2.0, 2.0, 100.0, 100.0],
25
+ }, size=200, w2v_cache_dir="weights/", load_pretrained=True):
26
+ super().__init__()
27
+
28
+ self.w2v_encoder = Wav2Vec2Encoder(cache_dir=w2v_cache_dir, load_pretrained=load_pretrained)
29
+ self.bridge = KANLinear(1024, 128)
30
+
31
+ self.d_args = d_args
32
+ filts = d_args["filts"]
33
+ gat_dims = d_args["gat_dims"]
34
+ pool_ratios = d_args["pool_ratios"]
35
+ temperatures = d_args["temperatures"]
36
+
37
+ self.first_bn = nn.BatchNorm2d(num_features=1)
38
+ self.selu = nn.SELU(inplace=True)
39
+ self.drop = nn.Dropout(0.5, inplace=True)
40
+ self.drop_way = nn.Dropout(0.2, inplace=True)
41
+
42
+ self.encoder = nn.Sequential(
43
+ nn.Sequential(Residual_block(nb_filts=filts[1], first=True)),
44
+ nn.Sequential(Residual_block(nb_filts=filts[2])),
45
+ nn.Sequential(Residual_block(nb_filts=filts[3])),
46
+ nn.Sequential(Residual_block(nb_filts=filts[4])),
47
+ nn.Sequential(Residual_block(nb_filts=filts[4])),
48
+ nn.Sequential(Residual_block(nb_filts=filts[4])))
49
+
50
+ self.pos_S = nn.Parameter(torch.randn(1, 42, filts[-1][-1]))
51
+ self.pos_T = nn.Parameter(torch.randn(1, 67, filts[-1][-1]))
52
+
53
+ self.GAT_layer_S = GraphAttentionLayer(filts[-1][-1], gat_dims[0], temperature=temperatures[0], size=size)
54
+ self.GAT_layer_T = GraphAttentionLayer(filts[-1][-1], gat_dims[0], temperature=temperatures[1], size=size)
55
+
56
+ self.pool_S = GraphPool(pool_ratios[0], gat_dims[0], 0.3, size=size)
57
+ self.pool_T = GraphPool(pool_ratios[1], gat_dims[0], 0.3, size=size)
58
+
59
+ self.master1 = nn.Parameter(torch.randn(1, 1, gat_dims[0]))
60
+ self.master2 = nn.Parameter(torch.randn(1, 1, gat_dims[0]))
61
+ self.master3 = nn.Parameter(torch.randn(1, 1, gat_dims[0]))
62
+ self.master4 = nn.Parameter(torch.randn(1, 1, gat_dims[0]))
63
+
64
+ self.inference_branch1 = InferenceBranch(
65
+ gat_dims=gat_dims,
66
+ temperature=temperatures[2],
67
+ pool_ratio=pool_ratios[2],
68
+ size=size
69
+ )
70
+ self.inference_branch2 = InferenceBranch(
71
+ gat_dims=gat_dims,
72
+ temperature=temperatures[2],
73
+ pool_ratio=pool_ratios[2],
74
+ size=size
75
+ )
76
+ self.inference_branch3 = InferenceBranch(
77
+ gat_dims=gat_dims,
78
+ temperature=temperatures[2],
79
+ pool_ratio=pool_ratios[2],
80
+ size=size
81
+ )
82
+ self.inference_branch4 = InferenceBranch(
83
+ gat_dims=gat_dims,
84
+ temperature=temperatures[2],
85
+ pool_ratio=pool_ratios[2],
86
+ size=size
87
+ )
88
+
89
+ self.out_layer = KANLinear(5 * gat_dims[1], 2)
90
+
91
+ def forward(self, x, Freq_aug=False):
92
+ x = self.w2v_encoder(x)
93
+ x = self.bridge(x)
94
+ x = x.transpose(1, 2)
95
+ x = x.unsqueeze(dim=1)
96
+ x = F.max_pool2d(torch.abs(x), (3, 3))
97
+ x = self.first_bn(x)
98
+ x = self.selu(x)
99
+
100
+ e = self.encoder(x)
101
+
102
+ # GAT-S
103
+ e_S, _ = torch.max(torch.abs(e), dim=3)
104
+ e_S = e_S.transpose(1, 2) + self.pos_S
105
+ gat_S = self.GAT_layer_S(e_S)
106
+ out_S = self.pool_S(gat_S)
107
+
108
+ # GAT-T
109
+ e_T, _ = torch.max(torch.abs(e), dim=2)
110
+ e_T = e_T.transpose(1, 2) + self.pos_T
111
+ gat_T = self.GAT_layer_T(e_T)
112
+ out_T = self.pool_T(gat_T)
113
+
114
+ out_T1, out_S1, master1 = self.inference_branch1(out_T, out_S, self.master1)
115
+ out_T2, out_S2, master2 = self.inference_branch2(out_T, out_S, self.master2)
116
+ out_T3, out_S3, master3 = self.inference_branch3(out_T, out_S, self.master3)
117
+ out_T4, out_S4, master4 = self.inference_branch4(out_T, out_S, self.master4)
118
+
119
+ out_T1, out_T2 = self.drop_way(out_T1), self.drop_way(out_T2)
120
+ out_T3, out_T4 = self.drop_way(out_T3), self.drop_way(out_T4)
121
+ out_S1, out_S2 = self.drop_way(out_S1), self.drop_way(out_S2)
122
+ out_S3, out_S4 = self.drop_way(out_S3), self.drop_way(out_S4)
123
+ master1, master2 = self.drop_way(master1), self.drop_way(master2)
124
+ master3, master4 = self.drop_way(master3), self.drop_way(master4)
125
+
126
+ out_T = torch.stack([out_T1, out_T2, out_T3, out_T4]).max(dim=0)[0]
127
+ out_S = torch.stack([out_S1, out_S2, out_S3, out_S4]).max(dim=0)[0]
128
+ master = torch.stack([master1, master2, master3, master4]).max(dim=0)[0]
129
+
130
+ T_max, _ = torch.max(torch.abs(out_T), dim=1)
131
+ T_avg = torch.mean(out_T, dim=1)
132
+ S_max, _ = torch.max(torch.abs(out_S), dim=1)
133
+ S_avg = torch.mean(out_S, dim=1)
134
+
135
+ last_hidden = torch.cat([T_max, T_avg, S_max, S_avg, master.squeeze(1)], dim=1)
136
+ last_hidden = self.drop(last_hidden)
137
+ output = self.out_layer(last_hidden)
138
+
139
+ return output
audio/aasist3/model/gat.py ADDED
@@ -0,0 +1,99 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, torch.nn as nn, torch.nn.functional as F
2
+
3
+ from .kan import KANLinear
4
+
5
+
6
+ class GraphAttentionLayer(nn.Module):
7
+ def __init__(self, in_dim, out_dim, size, **kwargs):
8
+ super().__init__()
9
+
10
+ # attention map
11
+ self.att_proj = KANLinear(in_dim, out_dim)
12
+ self.att_weight = self._init_new_params(out_dim, 1)
13
+
14
+ # project
15
+ self.proj_with_att = KANLinear(in_dim, out_dim)
16
+ self.proj_without_att = KANLinear(in_dim, out_dim)
17
+
18
+ # batch norm
19
+ self.bn = nn.BatchNorm1d(out_dim)
20
+
21
+ # dropout for inputs
22
+ self.input_drop = nn.Dropout(p=0.2)
23
+
24
+ # activate
25
+ self.act = nn.SELU(inplace=True)
26
+
27
+ # temperature
28
+ self.temp = 1.
29
+ if "temperature" in kwargs:
30
+ self.temp = kwargs["temperature"]
31
+
32
+ def forward(self, x):
33
+ '''
34
+ x :(#bs, #node, #dim)
35
+ '''
36
+ # apply input dropout
37
+ x = self.input_drop(x)
38
+
39
+ # derive attention map
40
+ att_map = self._derive_att_map(x)
41
+
42
+ # projection
43
+ x = self._project(x, att_map)
44
+
45
+ # apply batch norm
46
+ x = self._apply_BN(x)
47
+ x = self.act(x)
48
+ return x
49
+
50
+ def _pairwise_mul_nodes(self, x):
51
+ '''
52
+ Calculates pairwise multiplication of nodes.
53
+ - for attention map
54
+ x :(#bs, #node, #dim)
55
+ out_shape :(#bs, #node, #node, #dim)
56
+ '''
57
+
58
+ nb_nodes = x.size(1)
59
+ x = x.unsqueeze(2).expand(-1, -1, nb_nodes, -1)
60
+ x_mirror = x.transpose(1, 2)
61
+
62
+ return x * x_mirror
63
+
64
+ def _derive_att_map(self, x):
65
+ '''
66
+ x :(#bs, #node, #dim)
67
+ out_shape :(#bs, #node, #node, 1)
68
+ '''
69
+ att_map = self._pairwise_mul_nodes(x)
70
+ # size: (#bs, #node, #node, #dim_out)
71
+ att_map = torch.tanh(self.att_proj(att_map))
72
+ # size: (#bs, #node, #node, 1)
73
+ att_map = torch.matmul(att_map, self.att_weight)
74
+
75
+ # apply temperature
76
+ att_map = att_map / self.temp
77
+
78
+ att_map = F.softmax(att_map, dim=-2)
79
+
80
+ return att_map
81
+
82
+ def _project(self, x, att_map):
83
+ x1 = self.proj_with_att(torch.matmul(att_map.squeeze(-1), x))
84
+ x2 = self.proj_without_att(x)
85
+
86
+ return x1 + x2
87
+
88
+ def _apply_BN(self, x):
89
+ org_size = x.size()
90
+ x = x.view(-1, org_size[-1])
91
+ x = self.bn(x)
92
+ x = x.view(org_size)
93
+
94
+ return x
95
+
96
+ def _init_new_params(self, *size):
97
+ out = nn.Parameter(torch.FloatTensor(*size))
98
+ nn.init.xavier_normal_(out)
99
+ return out
audio/aasist3/model/hs_gal.py ADDED
@@ -0,0 +1,176 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, torch.nn as nn, torch.nn.functional as F
2
+
3
+ from .kan import KANLinear
4
+
5
+
6
+ class HtrgGraphAttentionLayer(nn.Module):
7
+ def __init__(self, in_dim, out_dim, size, **kwargs):
8
+ super().__init__()
9
+
10
+ self.proj_type1 = KANLinear(in_dim, in_dim)
11
+ self.proj_type2 = KANLinear(in_dim, in_dim)
12
+
13
+ # attention map
14
+ self.att_proj = KANLinear(in_dim, out_dim)
15
+ self.att_projM = KANLinear(in_dim, out_dim)
16
+
17
+ self.att_weight11 = self._init_new_params(out_dim, 1)
18
+ self.att_weight22 = self._init_new_params(out_dim, 1)
19
+ self.att_weight12 = self._init_new_params(out_dim, 1)
20
+ self.att_weightM = self._init_new_params(out_dim, 1)
21
+
22
+ # project
23
+ self.proj_with_att = KANLinear(in_dim, out_dim)
24
+ self.proj_without_att = KANLinear(in_dim, out_dim)
25
+
26
+ self.proj_with_attM = KANLinear(in_dim, out_dim)
27
+ self.proj_without_attM = KANLinear(in_dim, out_dim)
28
+
29
+ # batch norm
30
+ self.bn = nn.BatchNorm1d(out_dim)
31
+
32
+ # dropout for inputs
33
+ self.input_drop = nn.Dropout(p=0.2)
34
+
35
+ # activate
36
+ self.act = nn.SELU(inplace=True)
37
+
38
+ # temperature
39
+ self.temp = 1.
40
+ if "temperature" in kwargs:
41
+ self.temp = kwargs["temperature"]
42
+
43
+ def forward(self, x1, x2, master=None):
44
+ '''
45
+ x1 :(#bs, #node, #dim)
46
+ x2 :(#bs, #node, #dim)
47
+ '''
48
+ num_type1 = x1.size(1)
49
+ num_type2 = x2.size(1)
50
+
51
+ x1 = self.proj_type1(x1)
52
+ x2 = self.proj_type2(x2)
53
+
54
+ x = torch.cat([x1, x2], dim=1)
55
+
56
+ if master is None:
57
+ master = torch.mean(x, dim=1, keepdim=True)
58
+
59
+ # apply input dropout
60
+ x = self.input_drop(x)
61
+
62
+ # derive attention map
63
+ att_map = self._derive_att_map(x, num_type1, num_type2)
64
+
65
+ # directional edge for master node
66
+ master = self._update_master(x, master)
67
+
68
+ # projection
69
+ x = self._project(x, att_map)
70
+
71
+ # apply batch norm
72
+ x = self._apply_BN(x)
73
+ # x = self.act(x)
74
+
75
+ x1 = x.narrow(1, 0, num_type1)
76
+ x2 = x.narrow(1, num_type1, num_type2)
77
+
78
+ return x1, x2, master
79
+
80
+ def _update_master(self, x, master):
81
+
82
+ att_map = self._derive_att_map_master(x, master)
83
+ master = self._project_master(x, master, att_map)
84
+
85
+ return master
86
+
87
+ def _pairwise_mul_nodes(self, x):
88
+ '''
89
+ Calculates pairwise multiplication of nodes.
90
+ - for attention map
91
+ x :(#bs, #node, #dim)
92
+ out_shape :(#bs, #node, #node, #dim)
93
+ '''
94
+
95
+ nb_nodes = x.size(1)
96
+ x = x.unsqueeze(2).expand(-1, -1, nb_nodes, -1)
97
+ x_mirror = x.transpose(1, 2)
98
+
99
+ return x * x_mirror
100
+
101
+ def _derive_att_map_master(self, x, master):
102
+ '''
103
+ x :(#bs, #node, #dim)
104
+ out_shape :(#bs, #node, #node, 1)
105
+ '''
106
+ att_map = x * master
107
+ att_map = torch.tanh(self.att_projM(att_map))
108
+
109
+ att_map = torch.matmul(att_map, self.att_weightM)
110
+
111
+ # apply temperature
112
+ att_map = att_map / self.temp
113
+
114
+ att_map = F.softmax(att_map, dim=-2)
115
+
116
+ return att_map
117
+
118
+ def _derive_att_map(self, x, num_type1, num_type2):
119
+ '''
120
+ x :(#bs, #node, #dim)
121
+ out_shape :(#bs, #node, #node, 1)
122
+ '''
123
+ att_map = self._pairwise_mul_nodes(x)
124
+ # size: (#bs, #node, #node, #dim_out)
125
+ att_map = torch.tanh(self.att_proj(att_map))
126
+ # size: (#bs, #node, #node, 1)
127
+
128
+ att_board = torch.zeros_like(att_map[:, :, :, 0]).unsqueeze(-1)
129
+
130
+ att_board[:, :num_type1, :num_type1, :] = torch.matmul(
131
+ att_map[:, :num_type1, :num_type1, :], self.att_weight11)
132
+ att_board[:, num_type1:, num_type1:, :] = torch.matmul(
133
+ att_map[:, num_type1:, num_type1:, :], self.att_weight22)
134
+ att_board[:, :num_type1, num_type1:, :] = torch.matmul(
135
+ att_map[:, :num_type1, num_type1:, :], self.att_weight12)
136
+ att_board[:, num_type1:, :num_type1, :] = torch.matmul(
137
+ att_map[:, num_type1:, :num_type1, :], self.att_weight12)
138
+
139
+ att_map = att_board
140
+
141
+ # att_map = torch.matmul(att_map, self.att_weight12)
142
+
143
+ # apply temperature
144
+ att_map = att_map / self.temp
145
+
146
+ att_map = F.softmax(att_map, dim=-2)
147
+
148
+ return att_map
149
+
150
+ def _project(self, x, att_map):
151
+ x1 = self.proj_with_att(torch.matmul(att_map.squeeze(-1), x))
152
+ x2 = self.proj_without_att(x)
153
+
154
+ return x1 + x2
155
+
156
+ def _project_master(self, x, master, att_map):
157
+
158
+ x1 = self.proj_with_attM(torch.matmul(
159
+ att_map.squeeze(-1).unsqueeze(1), x))
160
+ x2 = self.proj_without_attM(master)
161
+
162
+ return x1 + x2
163
+
164
+ def _apply_BN(self, x):
165
+ org_size = x.size()
166
+ x = x.view(-1, org_size[-1])
167
+ x = self.bn(x)
168
+ x = x.view(org_size)
169
+
170
+ return x
171
+
172
+ def _init_new_params(self, *size):
173
+ out = nn.Parameter(torch.FloatTensor(*size))
174
+ nn.init.xavier_normal_(out)
175
+ return out
176
+
audio/aasist3/model/kan.py ADDED
@@ -0,0 +1,213 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, math, torch.nn.functional as F
2
+
3
+
4
+ class KANLinear(torch.nn.Module):
5
+ def __init__(
6
+ self,
7
+ in_features,
8
+ out_features,
9
+ grid_size=16,
10
+ spline_order=4,
11
+ scale_noise=0.1,
12
+ scale_base=1.0,
13
+ scale_spline=1.0,
14
+ enable_standalone_scale_spline=True,
15
+ base_activation=torch.nn.PReLU,
16
+ grid_eps=0.02,
17
+ grid_range=[-1, 1],
18
+ ):
19
+ super(KANLinear, self).__init__()
20
+ self.in_features = in_features
21
+ self.out_features = out_features
22
+ self.grid_size = grid_size
23
+ self.spline_order = spline_order
24
+
25
+ h = (grid_range[1] - grid_range[0]) / grid_size
26
+ grid = (
27
+ (
28
+ torch.arange(-spline_order, grid_size + spline_order + 1) * h
29
+ + grid_range[0]
30
+ )
31
+ .expand(in_features, -1)
32
+ .contiguous()
33
+ )
34
+ self.register_buffer("grid", grid)
35
+
36
+ self.base_weight = torch.nn.Parameter(torch.Tensor(out_features, in_features))
37
+ self.spline_weight = torch.nn.Parameter(
38
+ torch.Tensor(out_features, in_features, grid_size + spline_order)
39
+ )
40
+ if enable_standalone_scale_spline:
41
+ self.spline_scaler = torch.nn.Parameter(
42
+ torch.Tensor(out_features, in_features)
43
+ )
44
+
45
+ self.scale_noise = scale_noise
46
+ self.scale_base = scale_base
47
+ self.scale_spline = scale_spline
48
+ self.enable_standalone_scale_spline = enable_standalone_scale_spline
49
+ self.base_activation = base_activation()
50
+ self.grid_eps = grid_eps
51
+
52
+ self.reset_parameters()
53
+
54
+ def reset_parameters(self):
55
+ torch.nn.init.kaiming_uniform_(self.base_weight, a=math.sqrt(5) * self.scale_base)
56
+ with torch.no_grad():
57
+ noise = (
58
+ (
59
+ torch.rand(self.grid_size + 1, self.in_features, self.out_features)
60
+ - 1 / 2
61
+ )
62
+ * self.scale_noise
63
+ / self.grid_size
64
+ )
65
+ self.spline_weight.data.copy_(
66
+ (self.scale_spline if not self.enable_standalone_scale_spline else 1.0)
67
+ * self.curve2coeff(
68
+ self.grid.T[self.spline_order : -self.spline_order],
69
+ noise,
70
+ )
71
+ )
72
+ if self.enable_standalone_scale_spline:
73
+ # torch.nn.init.constant_(self.spline_scaler, self.scale_spline)
74
+ torch.nn.init.kaiming_uniform_(self.spline_scaler, a=math.sqrt(5) * self.scale_spline)
75
+
76
+ def b_splines(self, x: torch.Tensor):
77
+ """
78
+ Compute the B-spline bases for the given input tensor.
79
+
80
+ Args:
81
+ x (torch.Tensor): Input tensor of shape (batch_size, in_features).
82
+
83
+ Returns:
84
+ torch.Tensor: B-spline bases tensor of shape (batch_size, in_features, grid_size + spline_order).
85
+ """
86
+ assert x.dim() == 2 and x.size(1) == self.in_features
87
+
88
+ grid: torch.Tensor = (
89
+ self.grid
90
+ ) # (in_features, grid_size + 2 * spline_order + 1)
91
+ x = x.unsqueeze(-1)
92
+ bases = ((x >= grid[:, :-1]) & (x < grid[:, 1:])).to(x.dtype)
93
+ for k in range(1, self.spline_order + 1):
94
+ bases = (
95
+ (x - grid[:, : -(k + 1)])
96
+ / (grid[:, k:-1] - grid[:, : -(k + 1)])
97
+ * bases[:, :, :-1]
98
+ ) + (
99
+ (grid[:, k + 1 :] - x)
100
+ / (grid[:, k + 1 :] - grid[:, 1:(-k)])
101
+ * bases[:, :, 1:]
102
+ )
103
+
104
+ assert bases.size() == (
105
+ x.size(0),
106
+ self.in_features,
107
+ self.grid_size + self.spline_order,
108
+ )
109
+ return bases.contiguous()
110
+
111
+ def curve2coeff(self, x: torch.Tensor, y: torch.Tensor):
112
+ """
113
+ Compute the coefficients of the curve that interpolates the given points.
114
+
115
+ Args:
116
+ x (torch.Tensor): Input tensor of shape (batch_size, in_features).
117
+ y (torch.Tensor): Output tensor of shape (batch_size, in_features, out_features).
118
+
119
+ Returns:
120
+ torch.Tensor: Coefficients tensor of shape (out_features, in_features, grid_size + spline_order).
121
+ """
122
+ assert x.dim() == 2 and x.size(1) == self.in_features
123
+ assert y.size() == (x.size(0), self.in_features, self.out_features)
124
+
125
+ A = self.b_splines(x).transpose(
126
+ 0, 1
127
+ ) # (in_features, batch_size, grid_size + spline_order)
128
+ B = y.transpose(0, 1) # (in_features, batch_size, out_features)
129
+ solution = torch.linalg.lstsq(
130
+ A, B
131
+ ).solution # (in_features, grid_size + spline_order, out_features)
132
+ result = solution.permute(
133
+ 2, 0, 1
134
+ ) # (out_features, in_features, grid_size + spline_order)
135
+
136
+ assert result.size() == (
137
+ self.out_features,
138
+ self.in_features,
139
+ self.grid_size + self.spline_order,
140
+ )
141
+ return result.contiguous()
142
+
143
+ @property
144
+ def scaled_spline_weight(self):
145
+ return self.spline_weight * (
146
+ self.spline_scaler.unsqueeze(-1)
147
+ if self.enable_standalone_scale_spline
148
+ else 1.0
149
+ )
150
+
151
+ def forward(self, x: torch.Tensor):
152
+ assert x.size(-1) == self.in_features
153
+ original_shape = x.shape
154
+ x = x.reshape(-1, self.in_features)
155
+
156
+ base_output = F.linear(self.base_activation(x), self.base_weight)
157
+ spline_output = F.linear(
158
+ self.b_splines(x).view(x.size(0), -1),
159
+ self.scaled_spline_weight.reshape(self.out_features, -1),
160
+ )
161
+ output = base_output + spline_output
162
+ # print(*original_shape[:-1], output.shape)
163
+ output = output.view(*original_shape[:-1], self.out_features)
164
+ return output
165
+
166
+ @torch.no_grad()
167
+ def update_grid(self, x: torch.Tensor, margin=0.01):
168
+ assert x.dim() == 2 and x.size(1) == self.in_features
169
+ batch = x.size(0)
170
+
171
+ splines = self.b_splines(x) # (batch, in, coeff)
172
+ splines = splines.permute(1, 0, 2) # (in, batch, coeff)
173
+ orig_coeff = self.scaled_spline_weight # (out, in, coeff)
174
+ orig_coeff = orig_coeff.permute(1, 2, 0) # (in, coeff, out)
175
+ unreduced_spline_output = torch.bmm(splines, orig_coeff) # (in, batch, out)
176
+ unreduced_spline_output = unreduced_spline_output.permute(
177
+ 1, 0, 2
178
+ ) # (batch, in, out)
179
+
180
+ # sort each channel individually to collect data distribution
181
+ x_sorted = torch.sort(x, dim=0)[0]
182
+ grid_adaptive = x_sorted[
183
+ torch.linspace(
184
+ 0, batch - 1, self.grid_size + 1, dtype=torch.int64, device=x.device
185
+ )
186
+ ]
187
+
188
+ uniform_step = (x_sorted[-1] - x_sorted[0] + 2 * margin) / self.grid_size
189
+ grid_uniform = (
190
+ torch.arange(
191
+ self.grid_size + 1, dtype=torch.float32, device=x.device
192
+ ).unsqueeze(1)
193
+ * uniform_step
194
+ + x_sorted[0]
195
+ - margin
196
+ )
197
+
198
+ grid = self.grid_eps * grid_uniform + (1 - self.grid_eps) * grid_adaptive
199
+ grid = torch.concatenate(
200
+ [
201
+ grid[:1]
202
+ - uniform_step
203
+ * torch.arange(self.spline_order, 0, -1, device=x.device).unsqueeze(1),
204
+ grid,
205
+ grid[-1:]
206
+ + uniform_step
207
+ * torch.arange(1, self.spline_order + 1, device=x.device).unsqueeze(1),
208
+ ],
209
+ dim=0,
210
+ )
211
+
212
+ self.grid.copy_(grid.T)
213
+ self.spline_weight.data.copy_(self.curve2coeff(x, unreduced_spline_output))
audio/aasist3/model/pool.py ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, torch.nn as nn
2
+ from typing import Union
3
+
4
+ from .kan import KANLinear
5
+
6
+
7
+ class GraphPool(nn.Module):
8
+ def __init__(self, k: float, in_dim: int, p: Union[float, int], size):
9
+ super().__init__()
10
+ self.k = k
11
+ self.sigmoid = nn.Sigmoid()
12
+ self.proj = KANLinear(in_dim, 1)
13
+ self.drop = nn.Dropout(p=p) if p > 0 else nn.Identity()
14
+ self.in_dim = in_dim
15
+
16
+ def forward(self, h):
17
+ Z = self.drop(h)
18
+ weights = self.proj(Z)
19
+ scores = self.sigmoid(weights)
20
+ new_h = self.top_k_graph(scores, h, self.k)
21
+
22
+ return new_h
23
+
24
+ def top_k_graph(self, scores, h, k):
25
+ """
26
+ args
27
+ =====
28
+ scores: attention-based weights (#bs, #node, 1)
29
+ h: graph data (#bs, #node, #dim)
30
+ k: ratio of remaining nodes, (float)
31
+
32
+ returns
33
+ =====
34
+ h: graph pool applied data (#bs, #node', #dim)
35
+ """
36
+ _, n_nodes, n_feat = h.size()
37
+ n_nodes = max(int(n_nodes * k), 1)
38
+ _, idx = torch.topk(scores, n_nodes, dim=1)
39
+ idx = idx.expand(-1, -1, n_feat)
40
+
41
+ h = h * scores
42
+ h = torch.gather(h, 1, idx)
43
+
44
+ return h
45
+
audio/aasist3/model/residual.py ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch.nn as nn
2
+
3
+ class Residual_block(nn.Module):
4
+ def __init__(self, nb_filts, first=False):
5
+ super().__init__()
6
+ self.first = first
7
+
8
+ if not self.first:
9
+ self.bn1 = nn.BatchNorm2d(num_features=nb_filts[0])
10
+ self.conv1 = nn.Conv2d(in_channels=nb_filts[0],
11
+ out_channels=nb_filts[1],
12
+ kernel_size=(2, 3),
13
+ padding=(1, 1),
14
+ stride=1)
15
+ self.selu = nn.SELU()
16
+
17
+ self.bn2 = nn.BatchNorm2d(num_features=nb_filts[1])
18
+ self.conv2 = nn.Conv2d(in_channels=nb_filts[1],
19
+ out_channels=nb_filts[1],
20
+ kernel_size=(2, 3),
21
+ padding=(0, 1),
22
+ stride=1)
23
+
24
+ if nb_filts[0] != nb_filts[1]:
25
+ self.downsample = True
26
+ self.conv_downsample = nn.Conv2d(in_channels=nb_filts[0],
27
+ out_channels=nb_filts[1],
28
+ padding=(0, 1),
29
+ kernel_size=(1, 3),
30
+ stride=1)
31
+
32
+ else:
33
+ self.downsample = False
34
+ # self.mp = nn.MaxPool2d((1, 3)) # self.mp = nn.MaxPool2d((1,4))
35
+
36
+ def forward(self, x):
37
+ identity = x
38
+ if not self.first:
39
+ out = self.bn1(x)
40
+ out = self.selu(out)
41
+ else:
42
+ out = x
43
+ out = self.conv1(x)
44
+
45
+ # print('out',out.shape)
46
+ out = self.bn2(out)
47
+ out = self.selu(out)
48
+ # print('out',out.shape)
49
+ out = self.conv2(out)
50
+ #print('conv2 out',out.shape)
51
+ if self.downsample:
52
+ identity = self.conv_downsample(identity)
53
+
54
+ out += identity
55
+ # out = self.mp(out)
56
+ return out
audio/aasist3/model/wav2vec.py ADDED
@@ -0,0 +1,82 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch, torch.nn as nn
2
+ from transformers import Wav2Vec2Model, Wav2Vec2Config
3
+
4
+
5
+ class Wav2Vec2Encoder(nn.Module):
6
+ """SSL encoder based on Hugging Face's Wav2Vec2 model."""
7
+
8
+ def __init__(self,
9
+ model_name_or_path: str = "facebook/wav2vec2-large-xlsr-53",
10
+ ssl_out_dim: int = 768,
11
+ use_ssl_n_layers: int = None,
12
+ freeze_ssl_n_layers: int = 0,
13
+ output_attentions: bool = False,
14
+ output_hidden_states: bool = False,
15
+ normalize_waveform: bool = True,
16
+ cache_dir: str = "weights",
17
+ load_pretrained: bool = True):
18
+ """Initialize the Wav2Vec2 encoder.
19
+
20
+ Args:
21
+ model_name_or_path: HuggingFace model name or path to local model.
22
+ ssl_out_dim: Output dimension of the Wav2Vec2 encoder.
23
+ use_ssl_n_layers: Number of Wav2Vec2 layers to use. If None, use all layers.
24
+ freeze_ssl_n_layers: Number of Wav2Vec2 layers to freeze during training.
25
+ output_attentions: Whether to output attentions.
26
+ output_hidden_states: Whether to output hidden states.
27
+ normalize_waveform: Whether to normalize the waveform input.
28
+ cache_dir: Directory to cache pretrained models.
29
+ load_pretrained: Whether to load pretrained weights. If False, initializes with random weights.
30
+ """
31
+ super().__init__()
32
+
33
+ self.model_name_or_path = model_name_or_path
34
+ self.ssl_out_dim = ssl_out_dim
35
+ self.use_ssl_n_layers = use_ssl_n_layers
36
+ self.freeze_ssl_n_layers = freeze_ssl_n_layers
37
+ self.output_attentions = output_attentions
38
+ self.output_hidden_states = output_hidden_states
39
+ self.normalize_waveform = normalize_waveform
40
+
41
+ if load_pretrained:
42
+ self.model = Wav2Vec2Model.from_pretrained(model_name_or_path, cache_dir=cache_dir)
43
+ else:
44
+ config = Wav2Vec2Config.from_pretrained(
45
+ model_name_or_path,
46
+ cache_dir=cache_dir,
47
+ local_files_only=False
48
+ )
49
+ self.model = Wav2Vec2Model(config)
50
+ self.model.init_weights()
51
+
52
+ def forward(self, x):
53
+ """Forward pass through the Wav2Vec2 encoder.
54
+
55
+ Args:
56
+ x: Input tensor of shape (batch_size, sequence_length, channels)
57
+
58
+ Returns:
59
+ Extracted features of shape (batch_size, sequence_length, ssl_out_dim)
60
+ """
61
+ # Handle shape: convert (batch_size, sequence_length, channels) to (batch_size, sequence_length)
62
+ if x.ndim == 3:
63
+ x = x.squeeze(-1) # Remove channel dimension if present
64
+
65
+ if self.normalize_waveform:
66
+ x = x / (torch.max(torch.abs(x), dim=1, keepdim=True)[0] + 1e-8)
67
+
68
+ outputs = self.model(
69
+ x,
70
+ output_attentions=self.output_attentions,
71
+ output_hidden_states=self.output_hidden_states,
72
+ return_dict=True
73
+ )
74
+
75
+ last_hidden_state = outputs.last_hidden_state
76
+
77
+ if self.use_ssl_n_layers is not None and self.output_hidden_states and outputs.hidden_states is not None:
78
+ selected = outputs.hidden_states[-self.use_ssl_n_layers:]
79
+ last_hidden_state = torch.mean(torch.stack(selected, dim=0), dim=0)
80
+ del outputs
81
+
82
+ return last_hidden_state
audio/aasist3/requirements.txt ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ torch==2.5.1
2
+ torchaudio==2.5.1
3
+ transformers>=4.40.0
4
+ huggingface-hub>=0.20.0
5
+ safetensors>=0.4.0
6
+ numpy<2.0
7
+ soundfile
8
+ fastapi
9
+ uvicorn[standard]
10
+ pydantic
11
+ python-multipart
audio/nes2net/.gitignore ADDED
@@ -0,0 +1 @@
 
 
1
+ weights/
audio/nes2net/Dockerfile ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM nvidia/cuda:12.1.1-cudnn8-runtime-ubuntu22.04
2
+
3
+ ENV DEBIAN_FRONTEND=noninteractive
4
+ ENV PYTHONUNBUFFERED=1
5
+
6
+ RUN apt-get update && apt-get install -y --no-install-recommends \
7
+ python3 python3-pip python3-dev \
8
+ git ffmpeg libsndfile1 wget \
9
+ build-essential g++ \
10
+ && rm -rf /var/lib/apt/lists/*
11
+
12
+ RUN ln -sf /usr/bin/python3 /usr/bin/python
13
+
14
+ WORKDIR /app
15
+
16
+ # Install PyTorch with CUDA 12.1
17
+ RUN pip install --no-cache-dir \
18
+ torch==2.5.1 torchaudio==2.5.1 \
19
+ --index-url https://download.pytorch.org/whl/cu121
20
+
21
+ COPY requirements.txt .
22
+ RUN pip install --no-cache-dir -r requirements.txt
23
+
24
+ # Clone fairseq with patched C extensions (same as ShiftySpeech)
25
+ RUN git clone https://github.com/facebookresearch/fairseq.git /app/fairseq_repo && \
26
+ cd /app/fairseq_repo && \
27
+ git checkout a54021305d6b3c4c5959ac9395135f63202db8f1 && \
28
+ sed -i 's/ext_modules=extensions/ext_modules=[]/' setup.py && \
29
+ pip install --no-cache-dir --no-deps -e .
30
+
31
+ COPY model_scripts /app/model_scripts
32
+ COPY api.py .
33
+ RUN mkdir -p /app/weights
34
+ COPY weights/ /app/weights/
35
+
36
+ EXPOSE 8004
37
+
38
+ RUN adduser --disabled-password --gecos '' appuser
39
+ USER appuser
40
+
41
+ CMD ["python", "api.py"]
audio/nes2net/api.py ADDED
@@ -0,0 +1,315 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Nes2Net (XLSR + Nested Res2Net TDNN) Audio Deepfake Detection API.
2
+
3
+ Detects synthetic speech using the Nes2Net model architecture:
4
+ - Frontend: XLSR wav2vec 2.0 (Self-Supervised Learning)
5
+ - Backend: Nested Res2Net TDNN with SE modules
6
+
7
+ Reference: https://github.com/TianchiLiu/Nes2Net
8
+ """
9
+
10
+ import argparse
11
+ import base64
12
+ import io
13
+ import logging
14
+ import os
15
+ import platform
16
+ import sys
17
+ import time
18
+ import warnings
19
+ from typing import Optional
20
+
21
+ import librosa
22
+ import numpy as np
23
+ import torch
24
+ import uvicorn
25
+ from fastapi import FastAPI, HTTPException
26
+ from pydantic import BaseModel, Field
27
+
28
+ # Suppress deprecation warnings from fairseq/omegaconf compatibility
29
+ warnings.filterwarnings("ignore", category=DeprecationWarning)
30
+
31
+ # Monkey-patch omegaconf for fairseq compatibility (older fairseq
32
+ # expects is_primitive_type which was removed in newer omegaconf).
33
+ import omegaconf._utils as _omegaconf_utils
34
+
35
+ if not hasattr(_omegaconf_utils, "is_primitive_type"):
36
+ _omegaconf_utils.is_primitive_type = lambda t: t in (int, float, bool, str, bytes)
37
+
38
+ # Configure logging
39
+ logging.basicConfig(
40
+ level=logging.INFO,
41
+ format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
42
+ )
43
+ logger = logging.getLogger("nes2net_api")
44
+
45
+ # Add the model code to the path
46
+ if "/app" not in sys.path:
47
+ sys.path.insert(0, "/app")
48
+
49
+ # Import model class (deferred to allow path setup)
50
+ try:
51
+ from model_scripts.wav2vec2_Nes2Net_X import (
52
+ wav2vec2_Nes2Net_no_Res_w_allT as Nes2NetModel,
53
+ )
54
+ except ImportError as e:
55
+ logger.error(f"Failed to import Nes2Net model: {e}")
56
+ Nes2NetModel = None
57
+
58
+ # Constants
59
+ MODEL_NAME = "nes2net"
60
+ MODEL_ID = "nes2net_xlsr_itw_valaug"
61
+ WEIGHTS_PATH = "/app/weights/nes2net_itw_valaug.pt"
62
+
63
+
64
+ def _get_device():
65
+ """Select optimal device: MPS (Apple) > CUDA (NVIDIA) > CPU."""
66
+ override = os.environ.get("DEEPSAFE_DEVICE", "").strip().lower()
67
+ if override == "cpu":
68
+ return torch.device("cpu")
69
+ if override == "cuda" and torch.cuda.is_available():
70
+ return torch.device("cuda")
71
+ if (
72
+ override == "mps"
73
+ and hasattr(torch.backends, "mps")
74
+ and torch.backends.mps.is_available()
75
+ ):
76
+ return torch.device("mps")
77
+ if override:
78
+ pass # Invalid override, fall through to auto-detect
79
+ if (
80
+ platform.system() == "Darwin"
81
+ and hasattr(torch.backends, "mps")
82
+ and torch.backends.mps.is_available()
83
+ ):
84
+ return torch.device("mps")
85
+ if torch.cuda.is_available():
86
+ return torch.device("cuda")
87
+ return torch.device("cpu")
88
+
89
+
90
+ DEVICE = _get_device()
91
+
92
+ if DEVICE.type == "cuda":
93
+ torch.backends.cudnn.benchmark = True
94
+ torch.set_float32_matmul_precision("high")
95
+
96
+ if DEVICE.type == "cuda":
97
+ logger.info(
98
+ "Device: cuda (%s, %.1f GB VRAM)",
99
+ torch.cuda.get_device_name(0),
100
+ torch.cuda.get_device_properties(0).total_memory / 1024**3,
101
+ )
102
+ else:
103
+ logger.warning(
104
+ "Device: %s (no CUDA available -- check nvidia-container-toolkit)",
105
+ DEVICE,
106
+ )
107
+
108
+ SAMPLE_RATE = 16000
109
+ TARGET_SAMPLES = 64600 # ~4.04 seconds at 16kHz
110
+
111
+ # Global model instance
112
+ model = None
113
+
114
+
115
+ class AudioInput(BaseModel):
116
+ """Request schema for audio deepfake detection."""
117
+
118
+ audio_data: str = Field(
119
+ ..., description="Base64 encoded audio string (WAV/MP3/etc)"
120
+ )
121
+ threshold: Optional[float] = Field(
122
+ 0.5, ge=0.0, le=1.0, description="Classification threshold"
123
+ )
124
+
125
+
126
+ app = FastAPI(
127
+ title="Nes2Net Audio Deepfake Detection API",
128
+ description=(
129
+ "Service for detecting synthetic speech using the "
130
+ "Nes2Net model (XLSR wav2vec 2.0 + Nested Res2Net TDNN)."
131
+ ),
132
+ version="1.0.0",
133
+ )
134
+
135
+
136
+ def load_model():
137
+ """Load the Nes2Net model with fine-tuned weights.
138
+
139
+ Returns:
140
+ The loaded model, or None if loading fails.
141
+ """
142
+ global model
143
+ if model is not None:
144
+ return model
145
+
146
+ logger.info(f"Loading Nes2Net model onto {DEVICE}...")
147
+
148
+ if Nes2NetModel is None:
149
+ logger.error("Nes2Net model class not available.")
150
+ return None
151
+
152
+ if not os.path.exists(WEIGHTS_PATH):
153
+ logger.error(f"Model weights not found at {WEIGHTS_PATH}")
154
+ return None
155
+
156
+ try:
157
+ args = argparse.Namespace(
158
+ n_output_logits=2,
159
+ dilation=2,
160
+ pool_func="mean",
161
+ SE_ratio=[1],
162
+ Nes_ratio=[8, 8],
163
+ )
164
+ model = Nes2NetModel(args, str(DEVICE))
165
+
166
+ # Load fine-tuned weights
167
+ try:
168
+ state_dict = torch.load(
169
+ WEIGHTS_PATH,
170
+ map_location=DEVICE,
171
+ weights_only=False,
172
+ )
173
+ except TypeError:
174
+ state_dict = torch.load(WEIGHTS_PATH, map_location=DEVICE)
175
+
176
+ model.load_state_dict(state_dict)
177
+ model.to(DEVICE)
178
+ model.eval()
179
+
180
+ logger.info("Nes2Net model loaded successfully.")
181
+ return model
182
+ except Exception as e:
183
+ logger.exception(f"Failed to load Nes2Net model: {e}")
184
+ model = None
185
+ return None
186
+
187
+
188
+ @app.on_event("startup")
189
+ async def startup_event():
190
+ """Load model on service startup."""
191
+ load_model()
192
+
193
+
194
+ def _gpu_health_info() -> dict:
195
+ """Return GPU metrics for the health endpoint."""
196
+ if torch.cuda.is_available() and DEVICE.type == "cuda":
197
+ return {
198
+ "gpu_name": torch.cuda.get_device_name(0),
199
+ "vram_used_mb": round(torch.cuda.memory_allocated(0) / 1024**2),
200
+ "vram_total_mb": round(
201
+ torch.cuda.get_device_properties(0).total_memory / 1024**2
202
+ ),
203
+ }
204
+ return {}
205
+
206
+
207
+ @app.get("/health")
208
+ async def health():
209
+ """Health check endpoint."""
210
+ return {
211
+ "status": "healthy" if model is not None else "degraded",
212
+ "model": MODEL_NAME,
213
+ "model_id": MODEL_ID,
214
+ "device": str(DEVICE),
215
+ "weights_found": os.path.exists(WEIGHTS_PATH),
216
+ **_gpu_health_info(),
217
+ }
218
+
219
+
220
+ def preprocess_audio(audio_bytes: bytes) -> torch.Tensor:
221
+ """Preprocess audio for Nes2Net inference.
222
+
223
+ Loads audio, resamples to 16kHz mono, and pads/trims
224
+ to TARGET_SAMPLES using tiling (matching original training
225
+ preprocessing).
226
+
227
+ Args:
228
+ audio_bytes: Raw audio file bytes.
229
+
230
+ Returns:
231
+ Audio tensor of shape (1, TARGET_SAMPLES).
232
+
233
+ Raises:
234
+ ValueError: If audio preprocessing fails.
235
+ """
236
+ try:
237
+ logger.info("Starting audio preprocessing...")
238
+ audio, sr = librosa.load(io.BytesIO(audio_bytes), sr=SAMPLE_RATE, mono=True)
239
+ logger.info(f"Audio loaded. Length: {len(audio)} samples at {sr}Hz")
240
+
241
+ # Pad/trim to TARGET_SAMPLES using tiling
242
+ if len(audio) >= TARGET_SAMPLES:
243
+ audio = audio[:TARGET_SAMPLES]
244
+ else:
245
+ num_repeats = TARGET_SAMPLES // len(audio) + 1
246
+ audio = np.tile(audio, num_repeats)[:TARGET_SAMPLES]
247
+
248
+ logger.info(f"Audio padded/trimmed to {TARGET_SAMPLES} samples")
249
+
250
+ audio_tensor = torch.FloatTensor(audio).unsqueeze(0).to(DEVICE)
251
+ return audio_tensor
252
+ except Exception as e:
253
+ logger.error(f"Error preprocessing audio: {e}")
254
+ raise ValueError(f"Audio preprocessing failed: {str(e)}")
255
+
256
+
257
+ @app.post("/predict")
258
+ async def predict(input_data: AudioInput):
259
+ """Run deepfake detection on base64-encoded audio.
260
+
261
+ The model outputs 2 logits: [spoof_score, bonafide_score].
262
+ Class 0 = spoof (fake), Class 1 = bonafide (real).
263
+ The returned probability is the spoof/fake probability.
264
+ """
265
+ if model is None:
266
+ if load_model() is None:
267
+ raise HTTPException(status_code=503, detail="Model not loaded")
268
+
269
+ try:
270
+ start_time = time.time()
271
+ logger.info(
272
+ f"Prediction request. Data size: " f"{len(input_data.audio_data)} chars"
273
+ )
274
+
275
+ # Decode base64 audio
276
+ audio_bytes = base64.b64decode(input_data.audio_data)
277
+
278
+ # Preprocess
279
+ audio_tensor = preprocess_audio(audio_bytes)
280
+
281
+ # Inference
282
+ logger.info("Starting model inference...")
283
+ with torch.no_grad():
284
+ output = model(audio_tensor)
285
+
286
+ # output shape: [batch, 2]
287
+ # Index 0 = spoof logit, Index 1 = bonafide logit
288
+ probs = torch.softmax(output, dim=1)
289
+ prob_fake = probs[0, 0].item()
290
+
291
+ prediction = 1 if prob_fake >= input_data.threshold else 0
292
+ verdict = "fake" if prediction == 1 else "real"
293
+ inference_time = time.time() - start_time
294
+
295
+ logger.info(
296
+ f"Prediction: {verdict} (prob_fake={prob_fake:.4f}, "
297
+ f"time={inference_time:.3f}s)"
298
+ )
299
+
300
+ return {
301
+ "model": MODEL_NAME,
302
+ "probability": float(prob_fake),
303
+ "prediction": int(prediction),
304
+ "class": verdict,
305
+ "inference_time": float(inference_time),
306
+ }
307
+
308
+ except Exception as e:
309
+ logger.exception(f"Error during prediction: {e}")
310
+ raise HTTPException(status_code=500, detail=str(e))
311
+
312
+
313
+ if __name__ == "__main__":
314
+ port = int(os.environ.get("MODEL_PORT", 8004))
315
+ uvicorn.run(app, host="0.0.0.0", port=port)
audio/nes2net/model_scripts/__init__.py ADDED
File without changes
audio/nes2net/model_scripts/wav2vec2_Nes2Net_X.py ADDED
@@ -0,0 +1,317 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+
3
+ import fairseq
4
+ import torch
5
+ import torch.nn as nn
6
+
7
+ ___author__ = "Tianchi Liu"
8
+ __email__ = "tianchi_liu@u.nus.edu"
9
+ # modified from the model script from Hemlata Tak
10
+
11
+
12
+ class SSLModel(nn.Module):
13
+ def __init__(self, device):
14
+ super(SSLModel, self).__init__()
15
+ cp_path = (
16
+ "/app/weights/xlsr2_300m.pt" # Change the pre-trained XLSR model path.
17
+ )
18
+ model, cfg, task = fairseq.checkpoint_utils.load_model_ensemble_and_task(
19
+ [cp_path]
20
+ )
21
+ self.model = model[0]
22
+ self.device = device
23
+ self.out_dim = 1024
24
+ return
25
+
26
+ def extract_feat(self, input_data):
27
+ # put the model to GPU if it not there
28
+ if (
29
+ next(self.model.parameters()).device != input_data.device
30
+ or next(self.model.parameters()).dtype != input_data.dtype
31
+ ):
32
+ self.model.to(input_data.device, dtype=input_data.dtype)
33
+ self.model.train()
34
+ if True:
35
+ # input should be in shape (batch, length)
36
+ if input_data.ndim == 3:
37
+ input_tmp = input_data[:, :, 0]
38
+ else:
39
+ input_tmp = input_data
40
+ # [batch, length, dim]
41
+ emb = self.model(input_tmp, mask=False, features_only=True)["x"]
42
+ return emb
43
+
44
+
45
+ class SEModule(nn.Module):
46
+ def __init__(self, channels, SE_ratio=8):
47
+ super(SEModule, self).__init__()
48
+ self.se = nn.Sequential(
49
+ nn.AdaptiveAvgPool1d(1),
50
+ nn.Conv1d(channels, channels // SE_ratio, kernel_size=1, padding=0),
51
+ nn.ReLU(),
52
+ nn.Conv1d(channels // SE_ratio, channels, kernel_size=1, padding=0),
53
+ nn.Sigmoid(),
54
+ )
55
+
56
+ def forward(self, input):
57
+ x = self.se(input)
58
+ return input * x
59
+
60
+
61
+ class Bottle2neck(nn.Module):
62
+
63
+ def __init__(
64
+ self, inplanes, planes, kernel_size=None, dilation=None, scale=8, SE_ratio=8
65
+ ):
66
+ super(Bottle2neck, self).__init__()
67
+ width = int(math.floor(planes / scale))
68
+ self.conv1 = nn.Conv1d(inplanes, width * scale, kernel_size=1)
69
+ self.bn1 = nn.BatchNorm1d(width * scale)
70
+ self.nums = scale - 1
71
+ convs = []
72
+ bns = []
73
+ weighted_sum = []
74
+ num_pad = math.floor(kernel_size / 2) * dilation
75
+ for i in range(self.nums):
76
+ convs.append(
77
+ nn.Conv2d(
78
+ width,
79
+ width,
80
+ kernel_size=(kernel_size, 1),
81
+ dilation=(dilation, 1),
82
+ padding=(num_pad, 0),
83
+ )
84
+ )
85
+ bns.append(nn.BatchNorm2d(width))
86
+ initial_value = torch.ones(1, 1, 1, i + 2) * (1 / (i + 2))
87
+ weighted_sum.append(nn.Parameter(initial_value, requires_grad=True))
88
+ self.weighted_sum = nn.ParameterList(weighted_sum)
89
+ self.convs = nn.ModuleList(convs)
90
+ self.bns = nn.ModuleList(bns)
91
+ self.conv3 = nn.Conv1d(width * scale, planes, kernel_size=1)
92
+ self.bn3 = nn.BatchNorm1d(planes)
93
+ self.relu = nn.ReLU()
94
+ self.width = width
95
+ self.se = SEModule(planes, SE_ratio)
96
+
97
+ def forward(self, x):
98
+ residual = x
99
+ out = self.conv1(x)
100
+ out = self.relu(out)
101
+ out = self.bn1(out).unsqueeze(-1) # bz c T 1
102
+
103
+ spx = torch.split(out, self.width, 1)
104
+ sp = spx[self.nums]
105
+ for i in range(self.nums):
106
+ sp = torch.cat((sp, spx[i]), -1)
107
+
108
+ sp = self.bns[i](self.relu(self.convs[i](sp)))
109
+ sp_s = sp * self.weighted_sum[i]
110
+ sp_s = torch.sum(sp_s, dim=-1, keepdim=False)
111
+
112
+ if i == 0:
113
+ out = sp_s
114
+ else:
115
+ out = torch.cat((out, sp_s), 1)
116
+ out = torch.cat((out, spx[self.nums].squeeze(-1)), 1)
117
+ out = self.conv3(out)
118
+ out = self.relu(out)
119
+ out = self.bn3(out)
120
+ out = self.se(out)
121
+ out += residual
122
+ return out
123
+
124
+
125
+ class ASTP(nn.Module):
126
+ """Attentive statistics pooling: Channel- and context-dependent
127
+ statistics pooling, first used in ECAPA_TDNN.
128
+ """
129
+
130
+ def __init__(self, in_dim, bottleneck_dim=128, global_context_att=False):
131
+ super(ASTP, self).__init__()
132
+ self.global_context_att = global_context_att
133
+
134
+ # Use Conv1d with stride == 1 rather than Linear, then we don't
135
+ # need to transpose inputs.
136
+ if global_context_att:
137
+ self.linear1 = nn.Conv1d(
138
+ in_dim * 3, bottleneck_dim, kernel_size=1
139
+ ) # equals W and b in the paper
140
+ else:
141
+ self.linear1 = nn.Conv1d(
142
+ in_dim, bottleneck_dim, kernel_size=1
143
+ ) # equals W and b in the paper
144
+ self.linear2 = nn.Conv1d(
145
+ bottleneck_dim, in_dim, kernel_size=1
146
+ ) # equals V and k in the paper
147
+
148
+ def forward(self, x):
149
+ """
150
+ x: a 3-dimensional tensor in tdnn-based architecture (B,F,T)
151
+ or a 4-dimensional tensor in resnet architecture (B,C,F,T)
152
+ 0-dim: batch-dimension, last-dim: time-dimension (frame-dimension)
153
+ """
154
+ if len(x.shape) == 4:
155
+ x = x.reshape(x.shape[0], x.shape[1] * x.shape[2], x.shape[3])
156
+ assert len(x.shape) == 3
157
+
158
+ if self.global_context_att:
159
+ context_mean = torch.mean(x, dim=-1, keepdim=True).expand_as(x)
160
+ context_std = torch.sqrt(
161
+ torch.var(x, dim=-1, keepdim=True) + 1e-10
162
+ ).expand_as(x)
163
+ x_in = torch.cat((x, context_mean, context_std), dim=1)
164
+ else:
165
+ x_in = x
166
+
167
+ # DON'T use ReLU here! ReLU may be hard to converge.
168
+ alpha = torch.tanh(self.linear1(x_in)) # alpha = F.relu(self.linear1(x_in))
169
+ alpha = torch.softmax(self.linear2(alpha), dim=2)
170
+ mean = torch.sum(alpha * x, dim=2)
171
+ var = torch.sum(alpha * (x**2), dim=2) - mean**2
172
+ std = torch.sqrt(var.clamp(min=1e-10))
173
+ return torch.cat([mean, std], dim=1)
174
+
175
+
176
+ class Nested_Res2Net_TDNN(nn.Module):
177
+
178
+ def __init__(
179
+ self,
180
+ Nes_ratio=[8, 8],
181
+ input_channel=1024,
182
+ n_output_logits=2,
183
+ dilation=2,
184
+ pool_func="mean",
185
+ SE_ratio=[8],
186
+ ):
187
+
188
+ super(Nested_Res2Net_TDNN, self).__init__()
189
+ self.Nes_ratio = Nes_ratio[0]
190
+ assert input_channel % Nes_ratio[0] == 0
191
+ C = input_channel // Nes_ratio[0]
192
+ self.C = C
193
+ Build_in_Res2Nets = []
194
+ bns = []
195
+ for i in range(Nes_ratio[0] - 1):
196
+ Build_in_Res2Nets.append(
197
+ Bottle2neck(
198
+ C,
199
+ C,
200
+ kernel_size=3,
201
+ dilation=dilation,
202
+ scale=Nes_ratio[1],
203
+ SE_ratio=SE_ratio[0],
204
+ )
205
+ )
206
+ bns.append(nn.BatchNorm1d(C))
207
+ self.Build_in_Res2Nets = nn.ModuleList(Build_in_Res2Nets)
208
+ self.bns = nn.ModuleList(bns)
209
+ self.bn = nn.BatchNorm1d(1024)
210
+ self.relu = nn.ReLU()
211
+ self.pool_func = pool_func
212
+ if pool_func == "mean":
213
+ self.fc = nn.Linear(1024, n_output_logits)
214
+ elif pool_func == "ASTP":
215
+ self.pooling = ASTP(
216
+ in_dim=input_channel, bottleneck_dim=128, global_context_att=False
217
+ )
218
+ self.fc = nn.Linear(2048, n_output_logits)
219
+
220
+ def forward(self, x):
221
+ spx = torch.split(x, self.C, 1)
222
+ for i in range(self.Nes_ratio - 1):
223
+ if i == 0:
224
+ sp = spx[i]
225
+ else:
226
+ sp = sp + spx[i]
227
+ sp = self.Build_in_Res2Nets[i](sp)
228
+ sp = self.relu(sp)
229
+ sp = self.bns[i](sp)
230
+ if i == 0:
231
+ out = sp
232
+ else:
233
+ out = torch.cat((out, sp), 1)
234
+ out = torch.cat((out, spx[-1]), 1)
235
+ out = self.bn(out)
236
+ out = self.relu(out)
237
+ if self.pool_func == "mean":
238
+ out = torch.mean(out, dim=-1)
239
+ elif self.pool_func == "ASTP":
240
+ out = self.pooling(out)
241
+ out = self.fc(out)
242
+ return out
243
+
244
+
245
+ class wav2vec2_Nes2Net_no_Res_w_allT(nn.Module):
246
+ def __init__(self, args, device):
247
+ super().__init__()
248
+ self.device = device
249
+
250
+ self.n_output_logits = args.n_output_logits
251
+
252
+ ####
253
+ # create network wav2vec 2.0
254
+ ####
255
+ self.ssl_model = SSLModel(self.device)
256
+ self.Nested_Res2Net_TDNN = Nested_Res2Net_TDNN(
257
+ Nes_ratio=args.Nes_ratio,
258
+ input_channel=1024,
259
+ n_output_logits=self.n_output_logits,
260
+ dilation=args.dilation,
261
+ pool_func=args.pool_func,
262
+ SE_ratio=args.SE_ratio,
263
+ )
264
+
265
+ def forward(self, x):
266
+ # -------pre-trained Wav2vec model fine tunning ------------------------##
267
+ x_ssl_feat = self.ssl_model.extract_feat(x.squeeze(-1))
268
+ x_ssl_feat = x_ssl_feat.permute(0, 2, 1)
269
+ output = self.Nested_Res2Net_TDNN(x_ssl_feat)
270
+
271
+ return output
272
+
273
+
274
+ if __name__ == "__main__":
275
+ import argparse
276
+
277
+ parser = argparse.ArgumentParser()
278
+ parser.add_argument("--n_output_logits", type=int, default=2)
279
+ parser.add_argument("--dilation", type=int, default=2) # not important
280
+ parser.add_argument(
281
+ "--pool_func",
282
+ type=str,
283
+ default="mean",
284
+ choices=["mean", "ASTP"],
285
+ help="pooling function, choose from mean and ASTP",
286
+ )
287
+ parser.add_argument(
288
+ "--Nes_ratio",
289
+ type=int,
290
+ nargs="+",
291
+ default=[8, 8],
292
+ help="Nes_ratio, from outer to inner",
293
+ )
294
+ parser.add_argument(
295
+ "--SE_ratio",
296
+ type=int,
297
+ nargs="+",
298
+ default=[1],
299
+ help="SE downsampling ratio in the bottleneck",
300
+ )
301
+ args = parser.parse_args()
302
+
303
+ model = wav2vec2_Nes2Net_no_Res_w_allT(args=args, device="cpu")
304
+ x = torch.rand((4, 32000)).to("cpu")
305
+ model = model.to("cpu")
306
+ y = model(x)
307
+ print(y)
308
+ trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
309
+ print("all:", trainable_params)
310
+ trainable_params = sum(
311
+ p.numel() for p in model.ssl_model.parameters() if p.requires_grad
312
+ )
313
+ print("SSL:", trainable_params)
314
+ trainable_params = sum(
315
+ p.numel() for p in model.Nested_Res2Net_TDNN.parameters() if p.requires_grad
316
+ )
317
+ print("Backend:", trainable_params)
audio/nes2net/requirements.txt ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ torch==2.5.1
2
+ torchaudio==2.5.1
3
+ numpy==1.23.5
4
+ librosa==0.9.1
5
+ soundfile
6
+ scipy
7
+ omegaconf
8
+ hydra-core
9
+ bitarray
10
+ fastapi
11
+ uvicorn[standard]
12
+ pydantic
13
+ python-multipart
audio/safeear/Dockerfile ADDED
@@ -0,0 +1,49 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM nvidia/cuda:12.1.1-cudnn8-runtime-ubuntu22.04
2
+
3
+ ENV DEBIAN_FRONTEND=noninteractive
4
+ ENV PYTHONUNBUFFERED=1
5
+
6
+ RUN apt-get update && apt-get install -y --no-install-recommends \
7
+ python3 python3-pip python3-dev \
8
+ git ffmpeg libsndfile1 wget \
9
+ build-essential \
10
+ && rm -rf /var/lib/apt/lists/*
11
+
12
+ RUN ln -sf /usr/bin/python3 /usr/bin/python
13
+
14
+ WORKDIR /app
15
+
16
+ # Install PyTorch with CUDA 12.1
17
+ RUN pip install --no-cache-dir \
18
+ torch==2.5.1 torchaudio==2.5.1 \
19
+ --index-url https://download.pytorch.org/whl/cu121
20
+
21
+ COPY requirements.txt .
22
+ RUN pip install --no-cache-dir -r requirements.txt
23
+
24
+ # Clone SafeEar repository (for model code imports)
25
+ RUN git clone --depth 1 https://github.com/LetterLiGo/SafeEar.git /app/safeear_repo
26
+
27
+ # Install the fairseq fork with C extensions PATCHED OUT
28
+ # (same proven patch used by ShiftySpeech and Nes2Net --
29
+ # C extensions are not needed for checkpoint loading)
30
+ WORKDIR /app/safeear_repo/fairseq_ours
31
+ RUN sed -i 's/ext_modules=extensions/ext_modules=[]/' setup.py && \
32
+ pip install --no-cache-dir --no-deps -e .
33
+ WORKDIR /app
34
+
35
+ # Download model weights from HuggingFace
36
+ RUN mkdir -p /app/weights && \
37
+ wget -q -O /app/weights/SpeechTokenizer.pt \
38
+ "https://huggingface.co/TEC2004/SafeEar-ASV19-spoof-detection/resolve/main/SpeechTokenizer.pt" && \
39
+ wget -q -O /app/weights/model.ckpt \
40
+ "https://huggingface.co/TEC2004/SafeEar-ASV19-spoof-detection/resolve/main/model.ckpt"
41
+
42
+ COPY api.py .
43
+
44
+ EXPOSE 8002
45
+
46
+ RUN adduser --disabled-password --gecos '' appuser
47
+ USER appuser
48
+
49
+ CMD ["python", "api.py"]
audio/safeear/api.py ADDED
@@ -0,0 +1,321 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """SafeEar audio deepfake detection API service.
2
+
3
+ Uses the SafeEar content privacy-preserving model (CCS 2024) to detect
4
+ synthetic speech. Two-stage pipeline:
5
+ 1. SpeechTokenizer (neural audio codec) decouples acoustic features
6
+ 2. SafeEar1s (transformer classifier) detects spoofing from acoustic tokens
7
+
8
+ Weights: HuggingFace TEC2004/SafeEar-ASV19-spoof-detection
9
+ """
10
+
11
+ import base64
12
+ import logging
13
+ import os
14
+ import sys
15
+ import tempfile
16
+ import time
17
+ from typing import Optional
18
+
19
+ import uvicorn
20
+ from fastapi import FastAPI, HTTPException
21
+ from pydantic import BaseModel, Field
22
+
23
+ logging.basicConfig(
24
+ level=logging.INFO,
25
+ format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
26
+ )
27
+ logger = logging.getLogger("safeear_api")
28
+
29
+ import platform
30
+
31
+ import librosa
32
+ import numpy as np
33
+ import torch
34
+
35
+ # Add SafeEar repo to path for model imports
36
+ SAFEEAR_REPO_PATH = os.environ.get(
37
+ "SAFEEAR_REPO_PATH",
38
+ os.path.join(os.path.dirname(__file__), "safeear_repo"),
39
+ )
40
+ if SAFEEAR_REPO_PATH not in sys.path:
41
+ sys.path.insert(0, SAFEEAR_REPO_PATH)
42
+
43
+
44
+ def _get_device():
45
+ """Select optimal device: MPS (Apple) > CUDA (NVIDIA) > CPU."""
46
+ override = os.environ.get("DEEPSAFE_DEVICE", "").strip().lower()
47
+ if override == "cpu":
48
+ return torch.device("cpu")
49
+ if override == "cuda" and torch.cuda.is_available():
50
+ return torch.device("cuda")
51
+ if (
52
+ override == "mps"
53
+ and hasattr(torch.backends, "mps")
54
+ and torch.backends.mps.is_available()
55
+ ):
56
+ return torch.device("mps")
57
+ if (
58
+ platform.system() == "Darwin"
59
+ and hasattr(torch.backends, "mps")
60
+ and torch.backends.mps.is_available()
61
+ ):
62
+ return torch.device("mps")
63
+ if torch.cuda.is_available():
64
+ return torch.device("cuda")
65
+ return torch.device("cpu")
66
+
67
+
68
+ # Constants
69
+ MODEL_NAME = "safeear"
70
+ WEIGHTS_DIR = os.environ.get(
71
+ "WEIGHTS_DIR",
72
+ os.path.join(os.path.dirname(__file__), "weights"),
73
+ )
74
+ DEVICE = _get_device()
75
+
76
+ if DEVICE.type == "cuda":
77
+ torch.backends.cudnn.benchmark = True
78
+ torch.set_float32_matmul_precision("high")
79
+
80
+ if DEVICE.type == "cuda":
81
+ logger.info(
82
+ "Device: cuda (%s, %.1f GB VRAM)",
83
+ torch.cuda.get_device_name(0),
84
+ torch.cuda.get_device_properties(0).total_memory / 1024**3,
85
+ )
86
+ else:
87
+ logger.warning(
88
+ "Device: %s (no CUDA available -- check nvidia-container-toolkit)",
89
+ DEVICE,
90
+ )
91
+
92
+ SAMPLE_RATE = 16000
93
+ MAX_AUDIO_LENGTH = 64600 # ~4 seconds at 16kHz (ASVspoof standard)
94
+ SOFTMAX_TEMPERATURE = 5.0 # Calibration temperature for out-of-distribution data
95
+ NUM_INFERENCE_PASSES = 5 # Monte Carlo passes for stable predictions
96
+
97
+ # Global model instances
98
+ decouple_model = None
99
+ detect_model = None
100
+
101
+
102
+ class AudioInput(BaseModel):
103
+ """Schema for audio prediction requests."""
104
+
105
+ audio_data: str = Field(
106
+ ..., description="Base64 encoded audio string (WAV/MP3/etc)"
107
+ )
108
+ threshold: Optional[float] = Field(
109
+ 0.5, ge=0.0, le=1.0, description="Classification threshold"
110
+ )
111
+
112
+
113
+ app = FastAPI(
114
+ title="SafeEar Audio Deepfake Detection API",
115
+ description="Content privacy-preserving deepfake detection using SafeEar.",
116
+ version="1.0.0",
117
+ )
118
+
119
+
120
+ def load_models():
121
+ """Load both the decouple model (SpeechTokenizer) and detect model."""
122
+ global decouple_model, detect_model
123
+
124
+ if decouple_model is not None and detect_model is not None:
125
+ return True
126
+
127
+ speech_tokenizer_path = os.path.join(WEIGHTS_DIR, "SpeechTokenizer.pt")
128
+ checkpoint_path = os.path.join(WEIGHTS_DIR, "model.ckpt")
129
+
130
+ if not os.path.exists(speech_tokenizer_path):
131
+ logger.error(f"SpeechTokenizer weights not found: {speech_tokenizer_path}")
132
+ return False
133
+ if not os.path.exists(checkpoint_path):
134
+ logger.error(f"Model checkpoint not found: {checkpoint_path}")
135
+ return False
136
+
137
+ try:
138
+ # --- Load SpeechTokenizer (decouple model) ---
139
+ from safeear.models.decouple import SpeechTokenizer
140
+
141
+ logger.info("Loading SpeechTokenizer...")
142
+ decouple_model = SpeechTokenizer(
143
+ n_filters=64,
144
+ strides=[8, 5, 4, 2],
145
+ dimension=1024,
146
+ semantic_dimension=768,
147
+ bidirectional=True,
148
+ dilation_base=2,
149
+ residual_kernel_size=3,
150
+ n_residual_layers=1,
151
+ lstm_layers=2,
152
+ activation="ELU",
153
+ codebook_size=1024,
154
+ n_q=8,
155
+ sample_rate=16000,
156
+ )
157
+ st_state = torch.load(speech_tokenizer_path, map_location="cpu")
158
+ decouple_model.load_state_dict(st_state)
159
+ decouple_model.to(DEVICE)
160
+ decouple_model.eval()
161
+ logger.info("SpeechTokenizer loaded.")
162
+
163
+ # --- Load SafeEar1s (detect model) from Lightning checkpoint ---
164
+ from safeear.models.safeear import SafeEar1s, SE_Rawformer_front
165
+
166
+ logger.info("Loading SafeEar1s detect model...")
167
+ detect_model = SafeEar1s(
168
+ front=SE_Rawformer_front(),
169
+ embedding_dim=1024,
170
+ dropout_rate=0.1,
171
+ attention_dropout=0.1,
172
+ stochastic_depth=0.1,
173
+ num_layers=2,
174
+ num_heads=8,
175
+ num_classes=2,
176
+ positional_embedding="sine",
177
+ mlp_ratio=1.0,
178
+ )
179
+
180
+ # The .ckpt is a PyTorch Lightning checkpoint
181
+ ckpt = torch.load(checkpoint_path, map_location="cpu")
182
+ state_dict = ckpt.get("state_dict", ckpt)
183
+
184
+ # Lightning prefixes keys with "detect_model."
185
+ detect_state = {}
186
+ for k, v in state_dict.items():
187
+ if k.startswith("detect_model."):
188
+ detect_state[k.replace("detect_model.", "", 1)] = v
189
+
190
+ detect_model.load_state_dict(detect_state)
191
+ detect_model.to(DEVICE)
192
+ detect_model.eval()
193
+ logger.info("SafeEar1s detect model loaded.")
194
+ return True
195
+
196
+ except Exception as e:
197
+ logger.exception(f"Failed to load SafeEar models: {e}")
198
+ decouple_model = None
199
+ detect_model = None
200
+ return False
201
+
202
+
203
+ def preprocess_audio(audio_bytes: bytes) -> torch.Tensor:
204
+ """Load audio bytes, resample to 16kHz mono, pad/trim."""
205
+ with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
206
+ tmp.write(audio_bytes)
207
+ tmp_path = tmp.name
208
+
209
+ try:
210
+ waveform, _ = librosa.load(tmp_path, sr=SAMPLE_RATE, mono=True)
211
+ finally:
212
+ os.unlink(tmp_path)
213
+
214
+ if len(waveform) < MAX_AUDIO_LENGTH:
215
+ waveform = np.pad(waveform, (0, MAX_AUDIO_LENGTH - len(waveform)))
216
+ else:
217
+ waveform = waveform[:MAX_AUDIO_LENGTH]
218
+
219
+ # Shape: (1, 1, samples) -- batch=1, channels=1, time
220
+ tensor = torch.FloatTensor(waveform).unsqueeze(0).unsqueeze(0).to(DEVICE)
221
+ return tensor
222
+
223
+
224
+ @app.on_event("startup")
225
+ async def startup_event():
226
+ """Attempt to load models at startup."""
227
+ load_models()
228
+
229
+
230
+ def _gpu_health_info() -> dict:
231
+ """Return GPU metrics for the health endpoint."""
232
+ if torch.cuda.is_available() and DEVICE.type == "cuda":
233
+ return {
234
+ "gpu_name": torch.cuda.get_device_name(0),
235
+ "vram_used_mb": round(torch.cuda.memory_allocated(0) / 1024**2),
236
+ "vram_total_mb": round(
237
+ torch.cuda.get_device_properties(0).total_memory / 1024**2
238
+ ),
239
+ }
240
+ return {}
241
+
242
+
243
+ @app.get("/health")
244
+ async def health():
245
+ """Return service health status and model availability."""
246
+ models_loaded = decouple_model is not None and detect_model is not None
247
+ return {
248
+ "status": "healthy" if models_loaded else "degraded",
249
+ "model": MODEL_NAME,
250
+ "device": str(DEVICE),
251
+ "weights_found": (
252
+ os.path.exists(os.path.join(WEIGHTS_DIR, "SpeechTokenizer.pt"))
253
+ and os.path.exists(os.path.join(WEIGHTS_DIR, "model.ckpt"))
254
+ ),
255
+ **_gpu_health_info(),
256
+ }
257
+
258
+
259
+ @app.post("/predict")
260
+ async def predict(input_data: AudioInput):
261
+ """Run SafeEar inference on base64-encoded audio data."""
262
+ if decouple_model is None or detect_model is None:
263
+ if not load_models():
264
+ raise HTTPException(status_code=503, detail="Models not loaded")
265
+
266
+ try:
267
+ start_time = time.time()
268
+ logger.info(
269
+ "Received prediction request. "
270
+ f"Data size: {len(input_data.audio_data)} chars"
271
+ )
272
+
273
+ audio_bytes = base64.b64decode(input_data.audio_data)
274
+ x_wav = preprocess_audio(audio_bytes)
275
+
276
+ with torch.no_grad():
277
+ # Step 1: Extract acoustic tokens via SpeechTokenizer
278
+ # forward() returns:
279
+ # (reconstructed, commit_loss, semantic_feature, acoustic_tokens)
280
+ # layers=[0,1,2,3,4,5,6,7] means layer 0 goes to
281
+ # semantic_feature; layers 1-7 go to acoustic_tokens list
282
+ _, _, _, acoustic_tokens = decouple_model(
283
+ x_wav, layers=[0, 1, 2, 3, 4, 5, 6, 7]
284
+ )
285
+
286
+ # Step 2: Run detection model with Monte Carlo averaging
287
+ # SafeEar1s uses torch.randperm() in forward, so we average
288
+ # multiple passes for stable predictions
289
+ logit_sum = torch.zeros(1, 2, device=DEVICE)
290
+ for _ in range(NUM_INFERENCE_PASSES):
291
+ raw_logits, _ = detect_model(acoustic_tokens)
292
+ logit_sum += raw_logits
293
+ avg_logits = logit_sum / NUM_INFERENCE_PASSES
294
+
295
+ # Step 3: Get fake probability with temperature-scaled softmax
296
+ # The model produces extreme logits that saturate standard
297
+ # softmax. Temperature scaling preserves discrimination while
298
+ # giving more interpretable probabilities.
299
+ probs = torch.softmax(avg_logits / SOFTMAX_TEMPERATURE, dim=-1)
300
+ prob_fake = probs[0, 1].item()
301
+
302
+ prediction = 1 if prob_fake >= input_data.threshold else 0
303
+ verdict = "fake" if prediction == 1 else "real"
304
+ inference_time = time.time() - start_time
305
+
306
+ return {
307
+ "model": MODEL_NAME,
308
+ "probability": float(prob_fake),
309
+ "prediction": int(prediction),
310
+ "class": verdict,
311
+ "inference_time": float(inference_time),
312
+ }
313
+
314
+ except Exception as e:
315
+ logger.exception(f"Error during prediction: {e}")
316
+ raise HTTPException(status_code=500, detail=str(e))
317
+
318
+
319
+ if __name__ == "__main__":
320
+ port = int(os.environ.get("MODEL_PORT", 8002))
321
+ uvicorn.run(app, host="0.0.0.0", port=port)
audio/safeear/download_weights.sh ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ set -euo pipefail
3
+
4
+ WEIGHTS_DIR="${WEIGHTS_DIR:-/app/weights}"
5
+ REPO_DIR="${REPO_DIR:-/app/safeear_repo}"
6
+
7
+ mkdir -p "$WEIGHTS_DIR"
8
+
9
+ echo "==> Cloning SafeEar source repository..."
10
+ if [ ! -d "$REPO_DIR/.git" ]; then
11
+ git clone --depth 1 https://github.com/LetterLiGo/SafeEar.git "$REPO_DIR"
12
+ fi
13
+
14
+ echo "==> Downloading SpeechTokenizer.pt from HuggingFace..."
15
+ wget -q --show-progress -O "$WEIGHTS_DIR/SpeechTokenizer.pt" \
16
+ "https://huggingface.co/TEC2004/SafeEar-ASV19-spoof-detection/resolve/main/SpeechTokenizer.pt"
17
+
18
+ echo "==> Downloading model.ckpt from HuggingFace..."
19
+ wget -q --show-progress -O "$WEIGHTS_DIR/model.ckpt" \
20
+ "https://huggingface.co/TEC2004/SafeEar-ASV19-spoof-detection/resolve/main/model.ckpt"
21
+
22
+ echo "==> Weights downloaded to $WEIGHTS_DIR"
23
+ ls -lh "$WEIGHTS_DIR"
audio/safeear/requirements.txt ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ torch==2.5.1
2
+ torchaudio==2.5.1
3
+ librosa>=0.10.0
4
+ soundfile>=0.11.0
5
+ numpy>=1.23.0
6
+ einops>=0.7.0
7
+ timm>=0.9.0
8
+ hydra-core>=1.0.7
9
+ omegaconf>=2.1.0
10
+ pytorch-lightning>=1.6.0
11
+ scipy>=1.11.0
12
+ fastapi>=0.100.0
13
+ uvicorn[standard]>=0.20.0
14
+ python-multipart>=0.0.5
15
+ pydantic>=2.0.0
audio/shiftyspeech/Dockerfile ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM nvidia/cuda:12.1.1-cudnn8-runtime-ubuntu22.04
2
+
3
+ ENV DEBIAN_FRONTEND=noninteractive
4
+ ENV PYTHONUNBUFFERED=1
5
+
6
+ RUN apt-get update && apt-get install -y --no-install-recommends \
7
+ python3 python3-pip python3-dev \
8
+ git ffmpeg libsndfile1 \
9
+ build-essential g++ \
10
+ && rm -rf /var/lib/apt/lists/*
11
+
12
+ RUN ln -sf /usr/bin/python3 /usr/bin/python
13
+
14
+ WORKDIR /app
15
+
16
+ # Install PyTorch with CUDA 12.1
17
+ RUN pip install --no-cache-dir \
18
+ torch==2.5.1 torchaudio==2.5.1 \
19
+ --index-url https://download.pytorch.org/whl/cu121
20
+
21
+ COPY requirements.txt .
22
+ RUN pip install --no-cache-dir -r requirements.txt
23
+
24
+ # Clone fairseq with patched C extensions
25
+ RUN git clone https://github.com/facebookresearch/fairseq.git /app/fairseq_repo && \
26
+ cd /app/fairseq_repo && \
27
+ git checkout a54021305d6b3c4c5959ac9395135f63202db8f1 && \
28
+ sed -i 's/ext_modules=extensions/ext_modules=[]/' setup.py && \
29
+ pip install --no-cache-dir --no-deps -e .
30
+
31
+ COPY synthetic_speech_detection /app/synthetic_speech_detection
32
+ RUN mkdir -p /app/models
33
+ COPY api.py .
34
+ RUN mkdir -p /app/weights
35
+ COPY weights/ /app/weights/
36
+ RUN ln -sf /app/weights/xlsr_53_56k.pt /app/models/xlsr2_300m.pt
37
+
38
+ EXPOSE 8001
39
+
40
+ RUN adduser --disabled-password --gecos '' appuser
41
+ USER appuser
42
+
43
+ CMD ["python", "api.py"]
audio/shiftyspeech/api.py ADDED
@@ -0,0 +1,315 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """ShiftySpeech (SSL-AASIST) Audio Deepfake Detection API.
2
+
3
+ Detects synthetic speech using the SSL-AASIST model architecture:
4
+ - Frontend: XLSR wav2vec 2.0 (Self-Supervised Learning)
5
+ - Backend: AASIST (Audio Anti-Spoofing using Integrated
6
+ Spectro-Temporal Graph Attention Networks)
7
+
8
+ Reference: https://github.com/Ashigarg123/ShiftySpeech
9
+ """
10
+
11
+ import base64
12
+ import io
13
+ import logging
14
+ import os
15
+ import platform
16
+ import sys
17
+ import time
18
+ import warnings
19
+ from typing import Optional
20
+
21
+ import librosa
22
+ import numpy as np
23
+ import torch
24
+ import uvicorn
25
+ from fastapi import FastAPI, HTTPException
26
+ from pydantic import BaseModel, Field
27
+
28
+ # Suppress deprecation warnings from fairseq/omegaconf compatibility
29
+ warnings.filterwarnings("ignore", category=DeprecationWarning)
30
+
31
+ # Monkey-patch omegaconf for fairseq compatibility (older fairseq
32
+ # expects is_primitive_type which was removed in newer omegaconf).
33
+ import omegaconf._utils as _omegaconf_utils
34
+
35
+ if not hasattr(_omegaconf_utils, "is_primitive_type"):
36
+ _omegaconf_utils.is_primitive_type = lambda t: t in (int, float, bool, str, bytes)
37
+
38
+ # Configure logging
39
+ logging.basicConfig(
40
+ level=logging.INFO,
41
+ format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
42
+ )
43
+ logger = logging.getLogger("shiftyspeech_api")
44
+
45
+ # Add the SSL_Anti-spoofing model code to the path
46
+ MODEL_CODE_PATH = "/app/synthetic_speech_detection/SSL_Anti-spoofing"
47
+ if MODEL_CODE_PATH not in sys.path:
48
+ sys.path.insert(0, MODEL_CODE_PATH)
49
+
50
+ # Import model class (deferred to allow path setup)
51
+ try:
52
+ from model import Model as SSLAASISTModel
53
+ except ImportError as e:
54
+ logger.error(f"Failed to import SSL-AASIST model: {e}")
55
+ SSLAASISTModel = None
56
+
57
+ # Constants
58
+ MODEL_NAME = "shiftyspeech"
59
+ MODEL_ID = "ssl_aasist_augmented"
60
+ WEIGHTS_PATH = "/app/weights/hfg_aug_1_2.pt"
61
+ XLSR_DIR = "/app/models"
62
+
63
+
64
+ def _get_device():
65
+ """Select optimal device: MPS (Apple) > CUDA (NVIDIA) > CPU."""
66
+ override = os.environ.get("DEEPSAFE_DEVICE", "").strip().lower()
67
+ if override == "cpu":
68
+ return torch.device("cpu")
69
+ if override == "cuda" and torch.cuda.is_available():
70
+ return torch.device("cuda")
71
+ if (
72
+ override == "mps"
73
+ and hasattr(torch.backends, "mps")
74
+ and torch.backends.mps.is_available()
75
+ ):
76
+ return torch.device("mps")
77
+ if override:
78
+ pass # Invalid override, fall through to auto-detect
79
+ if (
80
+ platform.system() == "Darwin"
81
+ and hasattr(torch.backends, "mps")
82
+ and torch.backends.mps.is_available()
83
+ ):
84
+ return torch.device("mps")
85
+ if torch.cuda.is_available():
86
+ return torch.device("cuda")
87
+ return torch.device("cpu")
88
+
89
+
90
+ DEVICE = _get_device()
91
+
92
+ if DEVICE.type == "cuda":
93
+ torch.backends.cudnn.benchmark = True
94
+ torch.set_float32_matmul_precision("high")
95
+
96
+ if DEVICE.type == "cuda":
97
+ logger.info(
98
+ "Device: cuda (%s, %.1f GB VRAM)",
99
+ torch.cuda.get_device_name(0),
100
+ torch.cuda.get_device_properties(0).total_memory / 1024**3,
101
+ )
102
+ else:
103
+ logger.warning(
104
+ "Device: %s (no CUDA available -- check nvidia-container-toolkit)",
105
+ DEVICE,
106
+ )
107
+
108
+ SAMPLE_RATE = 16000
109
+ TARGET_SAMPLES = 64600 # ~4.04 seconds at 16kHz
110
+
111
+ # Global model instance
112
+ model = None
113
+
114
+
115
+ class AudioInput(BaseModel):
116
+ """Request schema for audio deepfake detection."""
117
+
118
+ audio_data: str = Field(
119
+ ..., description="Base64 encoded audio string (WAV/MP3/etc)"
120
+ )
121
+ threshold: Optional[float] = Field(
122
+ 0.5, ge=0.0, le=1.0, description="Classification threshold"
123
+ )
124
+
125
+
126
+ app = FastAPI(
127
+ title="ShiftySpeech Audio Deepfake Detection API",
128
+ description=(
129
+ "Service for detecting synthetic speech using the "
130
+ "SSL-AASIST model (XLSR wav2vec 2.0 + AASIST backend)."
131
+ ),
132
+ version="1.0.0",
133
+ )
134
+
135
+
136
+ def load_model():
137
+ """Load the SSL-AASIST model with augmented weights.
138
+
139
+ Returns:
140
+ The loaded model, or None if loading fails.
141
+ """
142
+ global model
143
+ if model is not None:
144
+ return model
145
+
146
+ logger.info(f"Loading SSL-AASIST model onto {DEVICE}...")
147
+
148
+ if SSLAASISTModel is None:
149
+ logger.error("SSL-AASIST model class not available.")
150
+ return None
151
+
152
+ if not os.path.exists(WEIGHTS_PATH):
153
+ logger.error(f"Model weights not found at {WEIGHTS_PATH}")
154
+ return None
155
+
156
+ try:
157
+ # Ensure XLSR model directory exists for architecture init
158
+ os.makedirs(XLSR_DIR, exist_ok=True)
159
+
160
+ import argparse
161
+
162
+ args = argparse.Namespace()
163
+ model = SSLAASISTModel(args, str(DEVICE))
164
+
165
+ # Load fine-tuned weights (includes XLSR weights)
166
+ try:
167
+ state_dict = torch.load(
168
+ WEIGHTS_PATH,
169
+ map_location=DEVICE,
170
+ weights_only=False,
171
+ )
172
+ except TypeError:
173
+ state_dict = torch.load(WEIGHTS_PATH, map_location=DEVICE)
174
+
175
+ model.load_state_dict(state_dict)
176
+ model.to(DEVICE)
177
+ model.eval()
178
+
179
+ logger.info("SSL-AASIST model loaded successfully.")
180
+ return model
181
+ except Exception as e:
182
+ logger.exception(f"Failed to load SSL-AASIST model: {e}")
183
+ model = None
184
+ return None
185
+
186
+
187
+ @app.on_event("startup")
188
+ async def startup_event():
189
+ """Load model on service startup."""
190
+ load_model()
191
+
192
+
193
+ def _gpu_health_info() -> dict:
194
+ """Return GPU metrics for the health endpoint."""
195
+ if torch.cuda.is_available() and DEVICE.type == "cuda":
196
+ return {
197
+ "gpu_name": torch.cuda.get_device_name(0),
198
+ "vram_used_mb": round(torch.cuda.memory_allocated(0) / 1024**2),
199
+ "vram_total_mb": round(
200
+ torch.cuda.get_device_properties(0).total_memory / 1024**2
201
+ ),
202
+ }
203
+ return {}
204
+
205
+
206
+ @app.get("/health")
207
+ async def health():
208
+ """Health check endpoint."""
209
+ return {
210
+ "status": "healthy" if model is not None else "degraded",
211
+ "model": MODEL_NAME,
212
+ "model_id": MODEL_ID,
213
+ "device": str(DEVICE),
214
+ "weights_found": os.path.exists(WEIGHTS_PATH),
215
+ **_gpu_health_info(),
216
+ }
217
+
218
+
219
+ def preprocess_audio(audio_bytes: bytes) -> torch.Tensor:
220
+ """Preprocess audio for SSL-AASIST inference.
221
+
222
+ Loads audio, resamples to 16kHz mono, and pads/trims
223
+ to TARGET_SAMPLES using tiling (matching original training
224
+ preprocessing from data_utils.py).
225
+
226
+ Args:
227
+ audio_bytes: Raw audio file bytes.
228
+
229
+ Returns:
230
+ Audio tensor of shape (1, TARGET_SAMPLES).
231
+
232
+ Raises:
233
+ ValueError: If audio preprocessing fails.
234
+ """
235
+ try:
236
+ logger.info("Starting audio preprocessing...")
237
+ audio, sr = librosa.load(io.BytesIO(audio_bytes), sr=SAMPLE_RATE, mono=True)
238
+ logger.info(f"Audio loaded. Length: {len(audio)} samples at {sr}Hz")
239
+
240
+ # Pad/trim to TARGET_SAMPLES using tiling
241
+ # (matches original data_utils.pad function)
242
+ if len(audio) >= TARGET_SAMPLES:
243
+ audio = audio[:TARGET_SAMPLES]
244
+ else:
245
+ num_repeats = TARGET_SAMPLES // len(audio) + 1
246
+ audio = np.tile(audio, num_repeats)[:TARGET_SAMPLES]
247
+
248
+ logger.info(f"Audio padded/trimmed to {TARGET_SAMPLES} samples")
249
+
250
+ audio_tensor = torch.FloatTensor(audio).unsqueeze(0).to(DEVICE)
251
+ return audio_tensor
252
+ except Exception as e:
253
+ logger.error(f"Error preprocessing audio: {e}")
254
+ raise ValueError(f"Audio preprocessing failed: {str(e)}")
255
+
256
+
257
+ @app.post("/predict")
258
+ async def predict(input_data: AudioInput):
259
+ """Run deepfake detection on base64-encoded audio.
260
+
261
+ The model outputs 2 logits: [spoof_score, bonafide_score].
262
+ Class 0 = spoof (fake), Class 1 = bonafide (real).
263
+ The returned probability is the spoof/fake probability.
264
+ """
265
+ if model is None:
266
+ if load_model() is None:
267
+ raise HTTPException(status_code=503, detail="Model not loaded")
268
+
269
+ try:
270
+ start_time = time.time()
271
+ logger.info(
272
+ f"Prediction request. Data size: " f"{len(input_data.audio_data)} chars"
273
+ )
274
+
275
+ # Decode base64 audio
276
+ audio_bytes = base64.b64decode(input_data.audio_data)
277
+
278
+ # Preprocess
279
+ audio_tensor = preprocess_audio(audio_bytes)
280
+
281
+ # Inference
282
+ logger.info("Starting model inference...")
283
+ with torch.no_grad():
284
+ output = model(audio_tensor)
285
+
286
+ # output shape: [batch, 2]
287
+ # Index 0 = spoof logit, Index 1 = bonafide logit
288
+ probs = torch.softmax(output, dim=1)
289
+ prob_fake = probs[0, 0].item()
290
+
291
+ prediction = 1 if prob_fake >= input_data.threshold else 0
292
+ verdict = "fake" if prediction == 1 else "real"
293
+ inference_time = time.time() - start_time
294
+
295
+ logger.info(
296
+ f"Prediction: {verdict} (prob_fake={prob_fake:.4f}, "
297
+ f"time={inference_time:.3f}s)"
298
+ )
299
+
300
+ return {
301
+ "model": MODEL_NAME,
302
+ "probability": float(prob_fake),
303
+ "prediction": int(prediction),
304
+ "class": verdict,
305
+ "inference_time": float(inference_time),
306
+ }
307
+
308
+ except Exception as e:
309
+ logger.exception(f"Error during prediction: {e}")
310
+ raise HTTPException(status_code=500, detail=str(e))
311
+
312
+
313
+ if __name__ == "__main__":
314
+ port = int(os.environ.get("MODEL_PORT", 8001))
315
+ uvicorn.run(app, host="0.0.0.0", port=port)
audio/shiftyspeech/evaluate.py ADDED
@@ -0,0 +1,210 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Evaluate ShiftySpeech SSL-AASIST on the DeepSafe audio dataset.
2
+
3
+ Reports accuracy, precision, recall, F1, EER, and per-file results.
4
+ """
5
+
6
+ import os
7
+ import sys
8
+ import time
9
+ import warnings
10
+
11
+ warnings.filterwarnings("ignore", category=DeprecationWarning)
12
+
13
+ # Monkey-patch omegaconf for fairseq compatibility
14
+ import omegaconf._utils as _omegaconf_utils
15
+
16
+ if not hasattr(_omegaconf_utils, "is_primitive_type"):
17
+ _omegaconf_utils.is_primitive_type = lambda t: t in (int, float, bool, str, bytes)
18
+
19
+ import argparse
20
+
21
+ import librosa
22
+ import numpy as np
23
+ import torch
24
+
25
+ # Add model code to path
26
+ SERVICE_DIR = os.path.dirname(os.path.abspath(__file__))
27
+ MODEL_CODE_PATH = os.path.join(
28
+ SERVICE_DIR, "synthetic_speech_detection", "SSL_Anti-spoofing"
29
+ )
30
+ sys.path.insert(0, MODEL_CODE_PATH)
31
+
32
+ from model import Model as SSLAASISTModel
33
+
34
+ SAMPLE_RATE = 16000
35
+ TARGET_SAMPLES = 64600
36
+ DATASET_DIR = os.path.join(
37
+ SERVICE_DIR, os.pardir, os.pardir, os.pardir, "dataset", "audio"
38
+ )
39
+ DATASET_DIR = os.path.normpath(DATASET_DIR)
40
+
41
+
42
+ def pad_audio(audio, target=TARGET_SAMPLES):
43
+ """Pad/trim audio to target length using tiling."""
44
+ if len(audio) >= target:
45
+ return audio[:target]
46
+ num_repeats = target // len(audio) + 1
47
+ return np.tile(audio, num_repeats)[:target]
48
+
49
+
50
+ def compute_eer(target_scores, nontarget_scores):
51
+ """Compute Equal Error Rate."""
52
+ n_scores = target_scores.size + nontarget_scores.size
53
+ all_scores = np.concatenate((target_scores, nontarget_scores))
54
+ labels = np.concatenate(
55
+ (np.ones(target_scores.size), np.zeros(nontarget_scores.size))
56
+ )
57
+ indices = np.argsort(all_scores, kind="mergesort")
58
+ labels = labels[indices]
59
+ tar_trial_sums = np.cumsum(labels)
60
+ nontarget_trial_sums = nontarget_scores.size - (
61
+ np.arange(1, n_scores + 1) - tar_trial_sums
62
+ )
63
+ frr = np.concatenate((np.atleast_1d(0), tar_trial_sums / target_scores.size))
64
+ far = np.concatenate(
65
+ (
66
+ np.atleast_1d(1),
67
+ nontarget_trial_sums / nontarget_scores.size,
68
+ )
69
+ )
70
+ abs_diffs = np.abs(frr - far)
71
+ min_index = np.argmin(abs_diffs)
72
+ eer = np.mean((frr[min_index], far[min_index]))
73
+ return eer
74
+
75
+
76
+ def main():
77
+ weights_path = os.path.join(SERVICE_DIR, "weights", "hfg_aug_1_2.pt")
78
+ os.makedirs(os.path.join(SERVICE_DIR, "models"), exist_ok=True)
79
+
80
+ print("=" * 70)
81
+ print("ShiftySpeech (SSL-AASIST) - DeepSafe Dataset Evaluation")
82
+ print("=" * 70)
83
+ print(f"Weights: {weights_path}")
84
+ print(f"Dataset: {DATASET_DIR}")
85
+ print(f"Device: cpu")
86
+ print()
87
+
88
+ # Load model
89
+ print("Loading model...")
90
+ start = time.time()
91
+ args_ns = argparse.Namespace()
92
+ ssl_model = SSLAASISTModel(args_ns, "cpu")
93
+ state_dict = torch.load(weights_path, map_location="cpu", weights_only=False)
94
+ ssl_model.load_state_dict(state_dict)
95
+ ssl_model.eval()
96
+ print(f"Model loaded in {time.time() - start:.1f}s")
97
+ print()
98
+
99
+ # Collect audio files
100
+ real_dir = os.path.join(DATASET_DIR, "real")
101
+ fake_dir = os.path.join(DATASET_DIR, "fake")
102
+
103
+ files = []
104
+ for fname in sorted(os.listdir(real_dir)):
105
+ if fname.endswith(".wav"):
106
+ files.append((os.path.join(real_dir, fname), 0, fname))
107
+ for fname in sorted(os.listdir(fake_dir)):
108
+ if fname.endswith(".wav"):
109
+ files.append((os.path.join(fake_dir, fname), 1, fname))
110
+
111
+ n_real = sum(1 for _, label, _ in files if label == 0)
112
+ n_fake = sum(1 for _, label, _ in files if label == 1)
113
+ print(f"Total files: {len(files)} (real: {n_real}, fake: {n_fake})")
114
+ print()
115
+
116
+ # Run inference
117
+ results = []
118
+ total_time = 0.0
119
+
120
+ print(
121
+ f"{'File':<20} {'True':>5} {'Pred':>5} {'P(fake)':>8} "
122
+ f"{'P(real)':>8} {'Time':>6}"
123
+ )
124
+ print("-" * 60)
125
+
126
+ for path, true_label, fname in files:
127
+ audio, sr = librosa.load(path, sr=SAMPLE_RATE, mono=True)
128
+ audio = pad_audio(audio)
129
+ x = torch.FloatTensor(audio).unsqueeze(0)
130
+
131
+ t0 = time.time()
132
+ with torch.no_grad():
133
+ out = ssl_model(x)
134
+ elapsed = time.time() - t0
135
+ total_time += elapsed
136
+
137
+ probs = torch.softmax(out, dim=1)
138
+ p_fake = probs[0, 0].item()
139
+ p_real = probs[0, 1].item()
140
+ pred = 1 if p_fake >= 0.5 else 0
141
+
142
+ results.append(
143
+ {
144
+ "file": fname,
145
+ "true_label": true_label,
146
+ "pred_label": pred,
147
+ "prob_fake": p_fake,
148
+ "prob_real": p_real,
149
+ }
150
+ )
151
+
152
+ true_str = "FAKE" if true_label == 1 else "REAL"
153
+ pred_str = "FAKE" if pred == 1 else "REAL"
154
+ correct = "ok" if pred == true_label else "XX"
155
+ print(
156
+ f"{fname:<20} {true_str:>5} {pred_str:>5} "
157
+ f"{p_fake:>8.4f} {p_real:>8.4f} {elapsed:>5.2f}s "
158
+ f"[{correct}]"
159
+ )
160
+
161
+ print()
162
+ print("=" * 70)
163
+ print("METRICS")
164
+ print("=" * 70)
165
+
166
+ # Compute metrics
167
+ true_labels = np.array([r["true_label"] for r in results])
168
+ pred_labels = np.array([r["pred_label"] for r in results])
169
+
170
+ tp = int(np.sum((pred_labels == 1) & (true_labels == 1)))
171
+ tn = int(np.sum((pred_labels == 0) & (true_labels == 0)))
172
+ fp = int(np.sum((pred_labels == 1) & (true_labels == 0)))
173
+ fn = int(np.sum((pred_labels == 0) & (true_labels == 1)))
174
+
175
+ accuracy = (tp + tn) / len(results) if len(results) > 0 else 0
176
+ precision = tp / (tp + fp) if (tp + fp) > 0 else 0
177
+ recall = tp / (tp + fn) if (tp + fn) > 0 else 0
178
+ f1 = (
179
+ 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0
180
+ )
181
+ specificity = tn / (tn + fp) if (tn + fp) > 0 else 0
182
+
183
+ # EER using bonafide scores (prob_real: higher = more real)
184
+ bonafide_scores = np.array(
185
+ [r["prob_real"] for r in results if r["true_label"] == 0]
186
+ )
187
+ spoof_scores = np.array([r["prob_real"] for r in results if r["true_label"] == 1])
188
+ if len(bonafide_scores) > 0 and len(spoof_scores) > 0:
189
+ eer = compute_eer(bonafide_scores, spoof_scores)
190
+ else:
191
+ eer = float("nan")
192
+
193
+ print(f"Accuracy: {accuracy:.4f} ({accuracy * 100:.1f}%)")
194
+ print(f"Precision: {precision:.4f}")
195
+ print(f"Recall: {recall:.4f}")
196
+ print(f"F1 Score: {f1:.4f}")
197
+ print(f"Specificity: {specificity:.4f}")
198
+ print(f"EER: {eer:.4f} ({eer * 100:.1f}%)")
199
+ print()
200
+ print(f"Confusion Matrix:")
201
+ print(f" TP={tp:>3d} FP={fp:>3d}")
202
+ print(f" FN={fn:>3d} TN={tn:>3d}")
203
+ print()
204
+ print(f"Total inference time: {total_time:.1f}s")
205
+ print(f"Avg per file: {total_time / len(results):.3f}s")
206
+ print(f"Total files: {len(results)}")
207
+
208
+
209
+ if __name__ == "__main__":
210
+ main()
audio/shiftyspeech/requirements.txt ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ fastapi
2
+ uvicorn
3
+ python-multipart
4
+ torch==2.5.1
5
+ torchaudio==2.5.1
6
+ numpy==1.23.5
7
+ scipy
8
+ librosa==0.9.1
9
+ soundfile
10
+ pydantic
11
+ omegaconf
12
+ hydra-core
13
+ scikit-learn
14
+ bitarray
audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/.env ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ WANDB_API_KEY="<wandb api-key>"
2
+ WANDB_PROJECT_NAME="SSL-AASIST"
audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2022 Hemlata
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/RawBoost.py ADDED
@@ -0,0 +1,143 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ # -*- coding: utf-8 -*-
3
+
4
+ import copy
5
+
6
+ import numpy as np
7
+ from scipy import signal
8
+
9
+ """
10
+ Hemlata Tak, Madhu Kamble, Jose Patino, Massimiliano Todisco, Nicholas Evans.
11
+ RawBoost: A Raw Data Boosting and Augmentation Method applied to Automatic Speaker Verification Anti-Spoofing.
12
+ In Proc. ICASSP 2022, pp:6382--6386.
13
+ """
14
+
15
+
16
+ def randRange(x1, x2, integer):
17
+ y = np.random.uniform(low=x1, high=x2, size=(1,))
18
+ if integer:
19
+ y = int(y)
20
+ return y
21
+
22
+
23
+ def normWav(x, always):
24
+ if always:
25
+ x = x / np.amax(abs(x))
26
+ elif np.amax(abs(x)) > 1:
27
+ x = x / np.amax(abs(x))
28
+ return x
29
+
30
+
31
+ def genNotchCoeffs(
32
+ nBands, minF, maxF, minBW, maxBW, minCoeff, maxCoeff, minG, maxG, fs
33
+ ):
34
+ b = 1
35
+ for i in range(0, nBands):
36
+ fc = randRange(minF, maxF, 0)
37
+ bw = randRange(minBW, maxBW, 0)
38
+ c = randRange(minCoeff, maxCoeff, 1)
39
+
40
+ if c / 2 == int(c / 2):
41
+ c = c + 1
42
+ f1 = fc - bw / 2
43
+ f2 = fc + bw / 2
44
+ if f1 <= 0:
45
+ f1 = 1 / 1000
46
+ if f2 >= fs / 2:
47
+ f2 = fs / 2 - 1 / 1000
48
+ b = np.convolve(
49
+ signal.firwin(c, [float(f1), float(f2)], window="hamming", fs=fs), b
50
+ )
51
+
52
+ G = randRange(minG, maxG, 0)
53
+ _, h = signal.freqz(b, 1, fs=fs)
54
+ b = pow(10, G / 20) * b / np.amax(abs(h))
55
+ return b
56
+
57
+
58
+ def filterFIR(x, b):
59
+ N = b.shape[0] + 1
60
+ xpad = np.pad(x, (0, N), "constant")
61
+ y = signal.lfilter(b, 1, xpad)
62
+ y = y[int(N / 2) : int(y.shape[0] - N / 2)]
63
+ return y
64
+
65
+
66
+ # Linear and non-linear convolutive noise
67
+ def LnL_convolutive_noise(
68
+ x,
69
+ N_f,
70
+ nBands,
71
+ minF,
72
+ maxF,
73
+ minBW,
74
+ maxBW,
75
+ minCoeff,
76
+ maxCoeff,
77
+ minG,
78
+ maxG,
79
+ minBiasLinNonLin,
80
+ maxBiasLinNonLin,
81
+ fs,
82
+ ):
83
+ y = [0] * x.shape[0]
84
+ for i in range(0, N_f):
85
+ if i == 1:
86
+ minG = minG - minBiasLinNonLin
87
+ maxG = maxG - maxBiasLinNonLin
88
+ b = genNotchCoeffs(
89
+ nBands, minF, maxF, minBW, maxBW, minCoeff, maxCoeff, minG, maxG, fs
90
+ )
91
+ y = y + filterFIR(np.power(x, (i + 1)), b)
92
+ y = y - np.mean(y)
93
+ y = normWav(y, 0)
94
+ return y
95
+
96
+
97
+ # Impulsive signal dependent noise
98
+ def ISD_additive_noise(x, P, g_sd):
99
+ beta = randRange(0, P, 0)
100
+
101
+ y = copy.deepcopy(x)
102
+ x_len = x.shape[0]
103
+ n = int(x_len * (beta / 100))
104
+ p = np.random.permutation(x_len)[:n]
105
+ f_r = np.multiply(
106
+ ((2 * np.random.rand(p.shape[0])) - 1), ((2 * np.random.rand(p.shape[0])) - 1)
107
+ )
108
+ r = g_sd * x[p] * f_r
109
+ y[p] = x[p] + r
110
+ y = normWav(y, 0)
111
+ return y
112
+
113
+
114
+ # Stationary signal independent noise
115
+
116
+
117
+ def SSI_additive_noise(
118
+ x,
119
+ SNRmin,
120
+ SNRmax,
121
+ nBands,
122
+ minF,
123
+ maxF,
124
+ minBW,
125
+ maxBW,
126
+ minCoeff,
127
+ maxCoeff,
128
+ minG,
129
+ maxG,
130
+ fs,
131
+ ):
132
+ noise = np.random.normal(0, 1, x.shape[0])
133
+ b = genNotchCoeffs(
134
+ nBands, minF, maxF, minBW, maxBW, minCoeff, maxCoeff, minG, maxG, fs
135
+ )
136
+ noise = filterFIR(noise, b)
137
+ noise = normWav(noise, 1)
138
+ SNR = randRange(SNRmin, SNRmax, 0)
139
+ noise = (
140
+ noise / np.linalg.norm(noise, 2) * np.linalg.norm(x, 2) / 10.0 ** (0.05 * SNR)
141
+ )
142
+ x = x + noise
143
+ return x
audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/Simplified_CM_solution.py ADDED
@@ -0,0 +1,227 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ from collections import OrderedDict
3
+
4
+ import fairseq
5
+ import numpy as np
6
+ import scipy.io as sio
7
+ import torch
8
+ import torch.nn as nn
9
+ import torch.nn.functional as F
10
+ from torch import Tensor
11
+ from torch.autograd import Variable
12
+ from torch.nn.parameter import Parameter
13
+ from torch.utils import data
14
+
15
+ ___author__ = "Hemlata Tak"
16
+ __email__ = "tak@eurecom.fr"
17
+
18
+ # from losses_anti_spoofing import AMSoftmax
19
+
20
+ ############################
21
+ ## FOR fine-tuning SSL MODEL
22
+ ############################
23
+
24
+
25
+ class SSLModel(nn.Module):
26
+ def __init__(self, device):
27
+ super(SSLModel, self).__init__()
28
+
29
+ cp_path = "/change_to_path_to_pre_trained_model_XLR_300M/xlsr2_300m.pt"
30
+ model, cfg, task = fairseq.checkpoint_utils.load_model_ensemble_and_task(
31
+ [cp_path]
32
+ )
33
+ self.model = model[0]
34
+ self.device = device
35
+ self.out_dim = 1024
36
+ return
37
+
38
+ def extract_feat(self, input_data):
39
+
40
+ # put the model to GPU if it not there
41
+ if (
42
+ next(self.model.parameters()).device != input_data.device
43
+ or next(self.model.parameters()).dtype != input_data.dtype
44
+ ):
45
+ self.model.to(input_data.device, dtype=input_data.dtype)
46
+ self.model.train()
47
+
48
+ if True:
49
+ # input should be in shape (batch, length)
50
+ if input_data.ndim == 3:
51
+ input_tmp = input_data[:, :, 0]
52
+ else:
53
+ input_tmp = input_data
54
+
55
+ # [batch, length, dim]
56
+ emb = self.model(input_tmp, mask=False, features_only=True)["x"]
57
+ return emb
58
+
59
+
60
+ # ---------Graph attention simple back-end------------------------#
61
+ """
62
+ Hemlata Tak, Jee-weon Jung, Jose Patino, Madhu Kamble, Massimiliano Todisco, Nicholas Evans.
63
+ End-to-end spectro-temporal graph attention networks for speaker verification anti-spoofing and speech deepfake detection.
64
+ In Proc. Automatic Speaker Verification and Spoofing Countermeasures Challenge 2021 Interspeech 2021 satellite workshop.
65
+ """
66
+
67
+
68
+ class GraphAttentionLayer(nn.Module):
69
+ def __init__(self, in_dim, out_dim, **kwargs):
70
+ super(GraphAttentionLayer, self).__init__()
71
+
72
+ # attention map
73
+ self.att_proj = nn.Linear(in_dim, out_dim)
74
+ self.att_weight = self._init_new_params(out_dim, 1)
75
+
76
+ # project
77
+ self.proj_with_att = nn.Linear(in_dim, out_dim)
78
+ self.proj_without_att = nn.Linear(in_dim, out_dim)
79
+
80
+ # batch norm
81
+ self.bn = nn.BatchNorm1d(out_dim)
82
+
83
+ # dropout for inputs
84
+ self.input_drop = nn.Dropout(p=0.2)
85
+
86
+ self.act = nn.SELU(inplace=True)
87
+
88
+ def forward(self, x):
89
+ """
90
+ x :(#bs, #node, #dim)
91
+ """
92
+ # apply input dropout
93
+ x = self.input_drop(x)
94
+
95
+ # derive attention map
96
+ att_map = self._derive_att_map(x)
97
+
98
+ # projection
99
+ x = self._project(x, att_map)
100
+
101
+ # apply batch norm
102
+ x = self._apply_BN(x)
103
+ x = self.act(x)
104
+
105
+ return x
106
+
107
+ def _pairwise_mul_nodes(self, x):
108
+ """
109
+ Calculates pairwise multiplication of nodes.
110
+ - for attention map
111
+ x :(#bs, #node, #dim)
112
+ out_shape :(#bs, #node, #node, #dim)
113
+ """
114
+
115
+ nb_nodes = x.size(1)
116
+ x = x.unsqueeze(2).expand(-1, -1, nb_nodes, -1)
117
+ x_mirror = x.transpose(1, 2)
118
+
119
+ return x * x_mirror
120
+
121
+ def _derive_att_map(self, x):
122
+ """
123
+ x :(#bs, #node, #dim)
124
+ out_shape :(#bs, #node, #node, 1)
125
+ """
126
+ att_map = self._pairwise_mul_nodes(x)
127
+ att_map = torch.tanh(
128
+ self.att_proj(att_map)
129
+ ) # size: (#bs, #node, #node, #dim_out)
130
+ att_map = torch.matmul(att_map, self.att_weight) # size: (#bs, #node, #node, 1)
131
+ att_map = F.softmax(att_map, dim=-2)
132
+
133
+ return att_map
134
+
135
+ def _project(self, x, att_map):
136
+ x1 = self.proj_with_att(torch.matmul(att_map.squeeze(-1), x))
137
+ x2 = self.proj_without_att(x)
138
+
139
+ return x1 + x2
140
+
141
+ def _apply_BN(self, x):
142
+ org_size = x.size()
143
+ x = x.view(-1, org_size[-1])
144
+ x = self.bn(x)
145
+ x = x.view(org_size)
146
+
147
+ return x
148
+
149
+ def _init_new_params(self, *size):
150
+ out = nn.Parameter(torch.FloatTensor(*size))
151
+ nn.init.xavier_normal_(out)
152
+ return out
153
+
154
+
155
+ class GraphPool(nn.Module):
156
+ def __init__(self, k: float, in_dim: int, p):
157
+ super().__init__()
158
+ self.k = k
159
+ self.sigmoid = nn.Sigmoid()
160
+ self.proj = nn.Linear(in_dim, 1)
161
+ self.drop = nn.Dropout(p=p) if p > 0 else nn.Identity()
162
+ self.in_dim = in_dim
163
+
164
+ def forward(self, h):
165
+ Z = self.drop(h)
166
+ weights = self.proj(Z)
167
+ scores = self.sigmoid(weights)
168
+ new_h = self.top_k_graph(scores, h, self.k)
169
+
170
+ return new_h
171
+
172
+ def top_k_graph(self, scores, h, k):
173
+ """
174
+ args
175
+ =====
176
+ scores: attention-based weights (#bs, #node, 1)
177
+ h: graph data (#bs, #node, #dim)
178
+ k: ratio of remaining nodes, (float)
179
+ returns
180
+ =====
181
+ h: graph pool applied data (#bs, #node', #dim)
182
+ """
183
+ _, n_nodes, n_feat = h.size()
184
+ n_nodes = max(int(n_nodes * k), 1)
185
+ _, idx = torch.topk(scores, n_nodes, dim=1)
186
+ idx = idx.expand(-1, -1, n_feat)
187
+
188
+ h = h * scores
189
+ h = torch.gather(h, 1, idx)
190
+
191
+ return h
192
+
193
+
194
+ class Model(nn.Module):
195
+ def __init__(self, d_args, device):
196
+ super(Model, self).__init__()
197
+
198
+ # SSL model
199
+ self.device = device
200
+ self.ssl_model = SSLModel(self.device)
201
+ self.LL = nn.Linear(self.ssl_model.out_dim, 128)
202
+ self.first_bn = nn.BatchNorm1d(num_features=128)
203
+ self.selu = nn.SELU(inplace=True)
204
+
205
+ # graph module layer
206
+ self.GAT_layer = GraphAttentionLayer(128, 64)
207
+ self.proj = nn.Linear(64, 1)
208
+ self.pool = GraphPool(0.8, 64, 0.3)
209
+
210
+ # classifier head
211
+ self.proj_node = nn.Linear(53, 2)
212
+
213
+ def forward(self, x_inp, Freq_aug=False):
214
+ # SSL wav2vec 2.0 model
215
+ x_ssl_feat = self.ssl_model.extract_feat(x_inp.squeeze(-1))
216
+ x_SSL = self.LL(x_ssl_feat) # (bs,frame_number,feat_out_dim)
217
+ x_SSL = x_SSL.transpose(1, 2) # (bs,feat_out_dim,frame_number)
218
+
219
+ x = F.max_pool1d(x_SSL, (3))
220
+ x = self.first_bn(x)
221
+ x = self.selu(x)
222
+
223
+ x = self.GAT_layer(x.transpose(1, 2))
224
+ x = self.pool(x)
225
+ x = self.proj(x).flatten(1)
226
+ output = self.proj_node(x)
227
+ return output
audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/data_utils.py ADDED
@@ -0,0 +1,292 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import random
3
+ from random import randrange
4
+
5
+ import librosa
6
+ import numpy as np
7
+ import torch
8
+ import torch.nn as nn
9
+ from RawBoost import (
10
+ ISD_additive_noise,
11
+ LnL_convolutive_noise,
12
+ SSI_additive_noise,
13
+ normWav,
14
+ )
15
+ from torch import Tensor
16
+ from torch.utils.data import Dataset
17
+
18
+ __author__ = "Hemlata Tak"
19
+ __email__ = "tak@eurecom.fr"
20
+
21
+
22
+ def genSpoof_list(dir_meta, is_train=False, is_eval=False):
23
+ d_meta = {}
24
+ file_list = []
25
+ with open(dir_meta, "r") as f:
26
+ l_meta = f.readlines()
27
+
28
+ if is_train:
29
+ for line in l_meta:
30
+ key, label = line.strip().split()
31
+ file_list.append(key)
32
+ d_meta[key] = 1 if label == "bonafide" else 0
33
+ return d_meta, file_list
34
+
35
+ elif is_eval:
36
+ for line in l_meta:
37
+ key, _ = line.strip().split(" ")
38
+ file_list.append(key)
39
+ return file_list
40
+ else:
41
+ for line in l_meta:
42
+ key, label = line.strip().split()
43
+ file_list.append(key)
44
+ d_meta[key] = 1 if label == "bonafide" else 0
45
+ return d_meta, file_list
46
+
47
+
48
+ def pad(x, max_len=64600):
49
+ x_len = x.shape[0]
50
+ if x_len >= max_len:
51
+ return x[:max_len]
52
+ # need to pad
53
+ num_repeats = int(max_len / x_len) + 1
54
+ padded_x = np.tile(x, (1, num_repeats))[:, :max_len][0]
55
+ return padded_x
56
+
57
+
58
+ class Dataset_ASVspoof2019_train(Dataset):
59
+ def __init__(self, args, metafile, algo):
60
+ """self.list_IDs : list of strings (each string: utt key),
61
+ self.labels: dictionary (key: utt key, value: label integer)"""
62
+
63
+ self.uttpath_labels = []
64
+ with open(metafile, "r") as f:
65
+ for line in f:
66
+ items = line.strip().split()
67
+ lb = 1 if items[-1] == "bonafide" else 0
68
+ self.uttpath_labels.append((items[0], lb))
69
+
70
+ self.algo = algo
71
+ self.args = args
72
+ self.cut = 64600 # take ~4 sec audio (64600 samples)
73
+
74
+ def __len__(self):
75
+ return len(self.uttpath_labels)
76
+
77
+ def __getitem__(self, index):
78
+ path, target = self.uttpath_labels[index]
79
+ X, fs = librosa.load(path, sr=16000)
80
+ Y = process_Rawboost_feature(X, fs, self.args, self.algo)
81
+ X_pad = pad(Y, self.cut)
82
+ x_inp = Tensor(X_pad)
83
+ return x_inp, target
84
+
85
+
86
+ class Dataset_ASVspoof2021_eval(Dataset):
87
+ def __init__(self, list_IDs):
88
+ """self.list_IDs : list of strings (each string: utt key),"""
89
+
90
+ self.list_IDs = list_IDs
91
+ self.cut = 64600 # take ~4 sec audio (64600 samples)
92
+
93
+ def __len__(self):
94
+ return len(self.list_IDs)
95
+
96
+ def __getitem__(self, index):
97
+ utt_id = self.list_IDs[index]
98
+ X, fs = librosa.load(utt_id, sr=16000)
99
+ X_pad = pad(X, self.cut)
100
+ x_inp = Tensor(X_pad)
101
+ return x_inp, utt_id
102
+
103
+
104
+ # --------------RawBoost data augmentation algorithms---------------------------##
105
+ def process_Rawboost_feature(feature, sr, args, algo):
106
+
107
+ # Data process by Convolutive noise (1st algo)
108
+ if algo == 1:
109
+
110
+ feature = LnL_convolutive_noise(
111
+ feature,
112
+ args.N_f,
113
+ args.nBands,
114
+ args.minF,
115
+ args.maxF,
116
+ args.minBW,
117
+ args.maxBW,
118
+ args.minCoeff,
119
+ args.maxCoeff,
120
+ args.minG,
121
+ args.maxG,
122
+ args.minBiasLinNonLin,
123
+ args.maxBiasLinNonLin,
124
+ sr,
125
+ )
126
+
127
+ # Data process by Impulsive noise (2nd algo)
128
+ elif algo == 2:
129
+
130
+ feature = ISD_additive_noise(feature, args.P, args.g_sd)
131
+
132
+ # Data process by coloured additive noise (3rd algo)
133
+ elif algo == 3:
134
+
135
+ feature = SSI_additive_noise(
136
+ feature,
137
+ args.SNRmin,
138
+ args.SNRmax,
139
+ args.nBands,
140
+ args.minF,
141
+ args.maxF,
142
+ args.minBW,
143
+ args.maxBW,
144
+ args.minCoeff,
145
+ args.maxCoeff,
146
+ args.minG,
147
+ args.maxG,
148
+ sr,
149
+ )
150
+
151
+ # Data process by all 3 algo. together in series (1+2+3)
152
+ elif algo == 4:
153
+
154
+ feature = LnL_convolutive_noise(
155
+ feature,
156
+ args.N_f,
157
+ args.nBands,
158
+ args.minF,
159
+ args.maxF,
160
+ args.minBW,
161
+ args.maxBW,
162
+ args.minCoeff,
163
+ args.maxCoeff,
164
+ args.minG,
165
+ args.maxG,
166
+ args.minBiasLinNonLin,
167
+ args.maxBiasLinNonLin,
168
+ sr,
169
+ )
170
+ feature = ISD_additive_noise(feature, args.P, args.g_sd)
171
+ feature = SSI_additive_noise(
172
+ feature,
173
+ args.SNRmin,
174
+ args.SNRmax,
175
+ args.nBands,
176
+ args.minF,
177
+ args.maxF,
178
+ args.minBW,
179
+ args.maxBW,
180
+ args.minCoeff,
181
+ args.maxCoeff,
182
+ args.minG,
183
+ args.maxG,
184
+ sr,
185
+ )
186
+
187
+ # Data process by 1st two algo. together in series (1+2)
188
+ elif algo == 5:
189
+
190
+ feature = LnL_convolutive_noise(
191
+ feature,
192
+ args.N_f,
193
+ args.nBands,
194
+ args.minF,
195
+ args.maxF,
196
+ args.minBW,
197
+ args.maxBW,
198
+ args.minCoeff,
199
+ args.maxCoeff,
200
+ args.minG,
201
+ args.maxG,
202
+ args.minBiasLinNonLin,
203
+ args.maxBiasLinNonLin,
204
+ sr,
205
+ )
206
+ feature = ISD_additive_noise(feature, args.P, args.g_sd)
207
+
208
+ # Data process by 1st and 3rd algo. together in series (1+3)
209
+ elif algo == 6:
210
+
211
+ feature = LnL_convolutive_noise(
212
+ feature,
213
+ args.N_f,
214
+ args.nBands,
215
+ args.minF,
216
+ args.maxF,
217
+ args.minBW,
218
+ args.maxBW,
219
+ args.minCoeff,
220
+ args.maxCoeff,
221
+ args.minG,
222
+ args.maxG,
223
+ args.minBiasLinNonLin,
224
+ args.maxBiasLinNonLin,
225
+ sr,
226
+ )
227
+ feature = SSI_additive_noise(
228
+ feature,
229
+ args.SNRmin,
230
+ args.SNRmax,
231
+ args.nBands,
232
+ args.minF,
233
+ args.maxF,
234
+ args.minBW,
235
+ args.maxBW,
236
+ args.minCoeff,
237
+ args.maxCoeff,
238
+ args.minG,
239
+ args.maxG,
240
+ sr,
241
+ )
242
+
243
+ # Data process by 2nd and 3rd algo. together in series (2+3)
244
+ elif algo == 7:
245
+
246
+ feature = ISD_additive_noise(feature, args.P, args.g_sd)
247
+ feature = SSI_additive_noise(
248
+ feature,
249
+ args.SNRmin,
250
+ args.SNRmax,
251
+ args.nBands,
252
+ args.minF,
253
+ args.maxF,
254
+ args.minBW,
255
+ args.maxBW,
256
+ args.minCoeff,
257
+ args.maxCoeff,
258
+ args.minG,
259
+ args.maxG,
260
+ sr,
261
+ )
262
+
263
+ # Data process by 1st two algo. together in Parallel (1||2)
264
+ elif algo == 8:
265
+
266
+ feature1 = LnL_convolutive_noise(
267
+ feature,
268
+ args.N_f,
269
+ args.nBands,
270
+ args.minF,
271
+ args.maxF,
272
+ args.minBW,
273
+ args.maxBW,
274
+ args.minCoeff,
275
+ args.maxCoeff,
276
+ args.minG,
277
+ args.maxG,
278
+ args.minBiasLinNonLin,
279
+ args.maxBiasLinNonLin,
280
+ sr,
281
+ )
282
+ feature2 = ISD_additive_noise(feature, args.P, args.g_sd)
283
+
284
+ feature_para = feature1 + feature2
285
+ feature = normWav(feature_para, 0) # normalized resultant waveform
286
+
287
+ # original data without Rawboost processing
288
+ else:
289
+
290
+ feature = feature
291
+
292
+ return feature
audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/model.py ADDED
@@ -0,0 +1,603 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import random
2
+ from typing import Union
3
+
4
+ import fairseq
5
+ import numpy as np
6
+ import torch
7
+ import torch.nn as nn
8
+ import torch.nn.functional as F
9
+ from torch import Tensor
10
+
11
+ ___author__ = "Hemlata Tak"
12
+ __email__ = "tak@eurecom.fr"
13
+
14
+ ############################
15
+ ## FOR fine-tuned SSL MODEL
16
+ ############################
17
+
18
+
19
+ class SSLModel(nn.Module):
20
+ def __init__(self, device):
21
+ super(SSLModel, self).__init__()
22
+
23
+ cp_path = "models/xlsr2_300m.pt" # Change the pre-trained XLSR model path.
24
+ model, cfg, task = fairseq.checkpoint_utils.load_model_ensemble_and_task(
25
+ [cp_path]
26
+ )
27
+ self.model = model[0]
28
+ self.device = device
29
+ self.out_dim = 1024
30
+ return
31
+
32
+ def extract_feat(self, input_data):
33
+
34
+ # put the model to GPU if it not there
35
+ if (
36
+ next(self.model.parameters()).device != input_data.device
37
+ or next(self.model.parameters()).dtype != input_data.dtype
38
+ ):
39
+ self.model.to(input_data.device, dtype=input_data.dtype)
40
+ self.model.train()
41
+
42
+ if True:
43
+ # input should be in shape (batch, length)
44
+ if input_data.ndim == 3:
45
+ input_tmp = input_data[:, :, 0]
46
+ else:
47
+ input_tmp = input_data
48
+
49
+ # [batch, length, dim]
50
+ emb = self.model(input_tmp, mask=False, features_only=True)["x"]
51
+ return emb
52
+
53
+
54
+ # ---------AASIST back-end------------------------#
55
+ """ Jee-weon Jung, Hee-Soo Heo, Hemlata Tak, Hye-jin Shim, Joon Son Chung, Bong-Jin Lee, Ha-Jin Yu and Nicholas Evans.
56
+ AASIST: Audio Anti-Spoofing Using Integrated Spectro-Temporal Graph Attention Networks.
57
+ In Proc. ICASSP 2022, pp: 6367--6371."""
58
+
59
+
60
+ class GraphAttentionLayer(nn.Module):
61
+ def __init__(self, in_dim, out_dim, **kwargs):
62
+ super().__init__()
63
+
64
+ # attention map
65
+ self.att_proj = nn.Linear(in_dim, out_dim)
66
+ self.att_weight = self._init_new_params(out_dim, 1)
67
+
68
+ # project
69
+ self.proj_with_att = nn.Linear(in_dim, out_dim)
70
+ self.proj_without_att = nn.Linear(in_dim, out_dim)
71
+
72
+ # batch norm
73
+ self.bn = nn.BatchNorm1d(out_dim)
74
+
75
+ # dropout for inputs
76
+ self.input_drop = nn.Dropout(p=0.2)
77
+
78
+ # activate
79
+ self.act = nn.SELU(inplace=True)
80
+
81
+ # temperature
82
+ self.temp = 1.0
83
+ if "temperature" in kwargs:
84
+ self.temp = kwargs["temperature"]
85
+
86
+ def forward(self, x):
87
+ """
88
+ x :(#bs, #node, #dim)
89
+ """
90
+ # apply input dropout
91
+ x = self.input_drop(x)
92
+
93
+ # derive attention map
94
+ att_map = self._derive_att_map(x)
95
+
96
+ # projection
97
+ x = self._project(x, att_map)
98
+
99
+ # apply batch norm
100
+ x = self._apply_BN(x)
101
+ x = self.act(x)
102
+ return x
103
+
104
+ def _pairwise_mul_nodes(self, x):
105
+ """
106
+ Calculates pairwise multiplication of nodes.
107
+ - for attention map
108
+ x :(#bs, #node, #dim)
109
+ out_shape :(#bs, #node, #node, #dim)
110
+ """
111
+
112
+ nb_nodes = x.size(1)
113
+ x = x.unsqueeze(2).expand(-1, -1, nb_nodes, -1)
114
+ x_mirror = x.transpose(1, 2)
115
+
116
+ return x * x_mirror
117
+
118
+ def _derive_att_map(self, x):
119
+ """
120
+ x :(#bs, #node, #dim)
121
+ out_shape :(#bs, #node, #node, 1)
122
+ """
123
+ att_map = self._pairwise_mul_nodes(x)
124
+ # size: (#bs, #node, #node, #dim_out)
125
+ att_map = torch.tanh(self.att_proj(att_map))
126
+ # size: (#bs, #node, #node, 1)
127
+ att_map = torch.matmul(att_map, self.att_weight)
128
+
129
+ # apply temperature
130
+ att_map = att_map / self.temp
131
+
132
+ att_map = F.softmax(att_map, dim=-2)
133
+
134
+ return att_map
135
+
136
+ def _project(self, x, att_map):
137
+ x1 = self.proj_with_att(torch.matmul(att_map.squeeze(-1), x))
138
+ x2 = self.proj_without_att(x)
139
+
140
+ return x1 + x2
141
+
142
+ def _apply_BN(self, x):
143
+ org_size = x.size()
144
+ x = x.view(-1, org_size[-1])
145
+ x = self.bn(x)
146
+ x = x.view(org_size)
147
+
148
+ return x
149
+
150
+ def _init_new_params(self, *size):
151
+ out = nn.Parameter(torch.FloatTensor(*size))
152
+ nn.init.xavier_normal_(out)
153
+ return out
154
+
155
+
156
+ class HtrgGraphAttentionLayer(nn.Module):
157
+ def __init__(self, in_dim, out_dim, **kwargs):
158
+ super().__init__()
159
+
160
+ self.proj_type1 = nn.Linear(in_dim, in_dim)
161
+ self.proj_type2 = nn.Linear(in_dim, in_dim)
162
+
163
+ # attention map
164
+ self.att_proj = nn.Linear(in_dim, out_dim)
165
+ self.att_projM = nn.Linear(in_dim, out_dim)
166
+
167
+ self.att_weight11 = self._init_new_params(out_dim, 1)
168
+ self.att_weight22 = self._init_new_params(out_dim, 1)
169
+ self.att_weight12 = self._init_new_params(out_dim, 1)
170
+ self.att_weightM = self._init_new_params(out_dim, 1)
171
+
172
+ # project
173
+ self.proj_with_att = nn.Linear(in_dim, out_dim)
174
+ self.proj_without_att = nn.Linear(in_dim, out_dim)
175
+
176
+ self.proj_with_attM = nn.Linear(in_dim, out_dim)
177
+ self.proj_without_attM = nn.Linear(in_dim, out_dim)
178
+
179
+ # batch norm
180
+ self.bn = nn.BatchNorm1d(out_dim)
181
+
182
+ # dropout for inputs
183
+ self.input_drop = nn.Dropout(p=0.2)
184
+
185
+ # activate
186
+ self.act = nn.SELU(inplace=True)
187
+
188
+ # temperature
189
+ self.temp = 1.0
190
+ if "temperature" in kwargs:
191
+ self.temp = kwargs["temperature"]
192
+
193
+ def forward(self, x1, x2, master=None):
194
+ """
195
+ x1 :(#bs, #node, #dim)
196
+ x2 :(#bs, #node, #dim)
197
+ """
198
+
199
+ num_type1 = x1.size(1)
200
+ num_type2 = x2.size(1)
201
+
202
+ x1 = self.proj_type1(x1)
203
+
204
+ x2 = self.proj_type2(x2)
205
+
206
+ x = torch.cat([x1, x2], dim=1)
207
+
208
+ if master is None:
209
+ master = torch.mean(x, dim=1, keepdim=True)
210
+
211
+ # apply input dropout
212
+ x = self.input_drop(x)
213
+
214
+ # derive attention map
215
+ att_map = self._derive_att_map(x, num_type1, num_type2)
216
+
217
+ # directional edge for master node
218
+ master = self._update_master(x, master)
219
+
220
+ # projection
221
+ x = self._project(x, att_map)
222
+
223
+ # apply batch norm
224
+ x = self._apply_BN(x)
225
+ x = self.act(x)
226
+
227
+ x1 = x.narrow(1, 0, num_type1)
228
+
229
+ x2 = x.narrow(1, num_type1, num_type2)
230
+
231
+ return x1, x2, master
232
+
233
+ def _update_master(self, x, master):
234
+
235
+ att_map = self._derive_att_map_master(x, master)
236
+ master = self._project_master(x, master, att_map)
237
+
238
+ return master
239
+
240
+ def _pairwise_mul_nodes(self, x):
241
+ """
242
+ Calculates pairwise multiplication of nodes.
243
+ - for attention map
244
+ x :(#bs, #node, #dim)
245
+ out_shape :(#bs, #node, #node, #dim)
246
+ """
247
+
248
+ nb_nodes = x.size(1)
249
+ x = x.unsqueeze(2).expand(-1, -1, nb_nodes, -1)
250
+ x_mirror = x.transpose(1, 2)
251
+
252
+ return x * x_mirror
253
+
254
+ def _derive_att_map_master(self, x, master):
255
+ """
256
+ x :(#bs, #node, #dim)
257
+ out_shape :(#bs, #node, #node, 1)
258
+ """
259
+ att_map = x * master
260
+ att_map = torch.tanh(self.att_projM(att_map))
261
+
262
+ att_map = torch.matmul(att_map, self.att_weightM)
263
+
264
+ # apply temperature
265
+ att_map = att_map / self.temp
266
+
267
+ att_map = F.softmax(att_map, dim=-2)
268
+
269
+ return att_map
270
+
271
+ def _derive_att_map(self, x, num_type1, num_type2):
272
+ """
273
+ x :(#bs, #node, #dim)
274
+ out_shape :(#bs, #node, #node, 1)
275
+ """
276
+ att_map = self._pairwise_mul_nodes(x)
277
+ # size: (#bs, #node, #node, #dim_out)
278
+ att_map = torch.tanh(self.att_proj(att_map))
279
+ # size: (#bs, #node, #node, 1)
280
+
281
+ att_board = torch.zeros_like(att_map[:, :, :, 0]).unsqueeze(-1)
282
+
283
+ att_board[:, :num_type1, :num_type1, :] = torch.matmul(
284
+ att_map[:, :num_type1, :num_type1, :], self.att_weight11
285
+ )
286
+ att_board[:, num_type1:, num_type1:, :] = torch.matmul(
287
+ att_map[:, num_type1:, num_type1:, :], self.att_weight22
288
+ )
289
+ att_board[:, :num_type1, num_type1:, :] = torch.matmul(
290
+ att_map[:, :num_type1, num_type1:, :], self.att_weight12
291
+ )
292
+ att_board[:, num_type1:, :num_type1, :] = torch.matmul(
293
+ att_map[:, num_type1:, :num_type1, :], self.att_weight12
294
+ )
295
+
296
+ att_map = att_board
297
+
298
+ # apply temperature
299
+ att_map = att_map / self.temp
300
+
301
+ att_map = F.softmax(att_map, dim=-2)
302
+
303
+ return att_map
304
+
305
+ def _project(self, x, att_map):
306
+ x1 = self.proj_with_att(torch.matmul(att_map.squeeze(-1), x))
307
+ x2 = self.proj_without_att(x)
308
+
309
+ return x1 + x2
310
+
311
+ def _project_master(self, x, master, att_map):
312
+
313
+ x1 = self.proj_with_attM(torch.matmul(att_map.squeeze(-1).unsqueeze(1), x))
314
+ x2 = self.proj_without_attM(master)
315
+
316
+ return x1 + x2
317
+
318
+ def _apply_BN(self, x):
319
+ org_size = x.size()
320
+ x = x.view(-1, org_size[-1])
321
+ x = self.bn(x)
322
+ x = x.view(org_size)
323
+
324
+ return x
325
+
326
+ def _init_new_params(self, *size):
327
+ out = nn.Parameter(torch.FloatTensor(*size))
328
+ nn.init.xavier_normal_(out)
329
+ return out
330
+
331
+
332
+ class GraphPool(nn.Module):
333
+ def __init__(self, k: float, in_dim: int, p: Union[float, int]):
334
+ super().__init__()
335
+ self.k = k
336
+ self.sigmoid = nn.Sigmoid()
337
+ self.proj = nn.Linear(in_dim, 1)
338
+ self.drop = nn.Dropout(p=p) if p > 0 else nn.Identity()
339
+ self.in_dim = in_dim
340
+
341
+ def forward(self, h):
342
+ Z = self.drop(h)
343
+ weights = self.proj(Z)
344
+ scores = self.sigmoid(weights)
345
+ new_h = self.top_k_graph(scores, h, self.k)
346
+
347
+ return new_h
348
+
349
+ def top_k_graph(self, scores, h, k):
350
+ """
351
+ args
352
+ =====
353
+ scores: attention-based weights (#bs, #node, 1)
354
+ h: graph data (#bs, #node, #dim)
355
+ k: ratio of remaining nodes, (float)
356
+ returns
357
+ =====
358
+ h: graph pool applied data (#bs, #node', #dim)
359
+ """
360
+ _, n_nodes, n_feat = h.size()
361
+ n_nodes = max(int(n_nodes * k), 1)
362
+ _, idx = torch.topk(scores, n_nodes, dim=1)
363
+ idx = idx.expand(-1, -1, n_feat)
364
+
365
+ h = h * scores
366
+ h = torch.gather(h, 1, idx)
367
+
368
+ return h
369
+
370
+
371
+ class Residual_block(nn.Module):
372
+ def __init__(self, nb_filts, first=False):
373
+ super().__init__()
374
+ self.first = first
375
+
376
+ if not self.first:
377
+ self.bn1 = nn.BatchNorm2d(num_features=nb_filts[0])
378
+ self.conv1 = nn.Conv2d(
379
+ in_channels=nb_filts[0],
380
+ out_channels=nb_filts[1],
381
+ kernel_size=(2, 3),
382
+ padding=(1, 1),
383
+ stride=1,
384
+ )
385
+ self.selu = nn.SELU(inplace=True)
386
+
387
+ self.bn2 = nn.BatchNorm2d(num_features=nb_filts[1])
388
+ self.conv2 = nn.Conv2d(
389
+ in_channels=nb_filts[1],
390
+ out_channels=nb_filts[1],
391
+ kernel_size=(2, 3),
392
+ padding=(0, 1),
393
+ stride=1,
394
+ )
395
+
396
+ if nb_filts[0] != nb_filts[1]:
397
+ self.downsample = True
398
+ self.conv_downsample = nn.Conv2d(
399
+ in_channels=nb_filts[0],
400
+ out_channels=nb_filts[1],
401
+ padding=(0, 1),
402
+ kernel_size=(1, 3),
403
+ stride=1,
404
+ )
405
+
406
+ else:
407
+ self.downsample = False
408
+
409
+ def forward(self, x):
410
+ identity = x
411
+ if not self.first:
412
+ out = self.bn1(x)
413
+ out = self.selu(out)
414
+ else:
415
+ out = x
416
+
417
+ out = self.conv1(x)
418
+
419
+ out = self.bn2(out)
420
+ out = self.selu(out)
421
+
422
+ out = self.conv2(out)
423
+
424
+ if self.downsample:
425
+ identity = self.conv_downsample(identity)
426
+
427
+ out += identity
428
+
429
+ return out
430
+
431
+
432
+ class Model(nn.Module):
433
+ def __init__(self, args, device):
434
+ super().__init__()
435
+ self.device = device
436
+
437
+ # AASIST parameters
438
+ filts = [128, [1, 32], [32, 32], [32, 64], [64, 64]]
439
+ gat_dims = [64, 32]
440
+ pool_ratios = [0.5, 0.5, 0.5, 0.5]
441
+ temperatures = [2.0, 2.0, 100.0, 100.0]
442
+
443
+ ####
444
+ # create network wav2vec 2.0
445
+ ####
446
+ self.ssl_model = SSLModel(self.device)
447
+ self.LL = nn.Linear(self.ssl_model.out_dim, 128)
448
+
449
+ self.first_bn = nn.BatchNorm2d(num_features=1)
450
+ self.first_bn1 = nn.BatchNorm2d(num_features=64)
451
+ self.drop = nn.Dropout(0.5, inplace=True)
452
+ self.drop_way = nn.Dropout(0.2, inplace=True)
453
+ self.selu = nn.SELU(inplace=True)
454
+
455
+ # RawNet2 encoder
456
+ self.encoder = nn.Sequential(
457
+ nn.Sequential(Residual_block(nb_filts=filts[1], first=True)),
458
+ nn.Sequential(Residual_block(nb_filts=filts[2])),
459
+ nn.Sequential(Residual_block(nb_filts=filts[3])),
460
+ nn.Sequential(Residual_block(nb_filts=filts[4])),
461
+ nn.Sequential(Residual_block(nb_filts=filts[4])),
462
+ nn.Sequential(Residual_block(nb_filts=filts[4])),
463
+ )
464
+
465
+ self.attention = nn.Sequential(
466
+ nn.Conv2d(64, 128, kernel_size=(1, 1)),
467
+ nn.SELU(inplace=True),
468
+ nn.BatchNorm2d(128),
469
+ nn.Conv2d(128, 64, kernel_size=(1, 1)),
470
+ )
471
+ # position encoding
472
+ self.pos_S = nn.Parameter(torch.randn(1, 42, filts[-1][-1]))
473
+
474
+ self.master1 = nn.Parameter(torch.randn(1, 1, gat_dims[0]))
475
+ self.master2 = nn.Parameter(torch.randn(1, 1, gat_dims[0]))
476
+
477
+ # Graph module
478
+ self.GAT_layer_S = GraphAttentionLayer(
479
+ filts[-1][-1], gat_dims[0], temperature=temperatures[0]
480
+ )
481
+ self.GAT_layer_T = GraphAttentionLayer(
482
+ filts[-1][-1], gat_dims[0], temperature=temperatures[1]
483
+ )
484
+ # HS-GAL layer
485
+ self.HtrgGAT_layer_ST11 = HtrgGraphAttentionLayer(
486
+ gat_dims[0], gat_dims[1], temperature=temperatures[2]
487
+ )
488
+ self.HtrgGAT_layer_ST12 = HtrgGraphAttentionLayer(
489
+ gat_dims[1], gat_dims[1], temperature=temperatures[2]
490
+ )
491
+ self.HtrgGAT_layer_ST21 = HtrgGraphAttentionLayer(
492
+ gat_dims[0], gat_dims[1], temperature=temperatures[2]
493
+ )
494
+ self.HtrgGAT_layer_ST22 = HtrgGraphAttentionLayer(
495
+ gat_dims[1], gat_dims[1], temperature=temperatures[2]
496
+ )
497
+
498
+ # Graph pooling layers
499
+ self.pool_S = GraphPool(pool_ratios[0], gat_dims[0], 0.3)
500
+ self.pool_T = GraphPool(pool_ratios[1], gat_dims[0], 0.3)
501
+ self.pool_hS1 = GraphPool(pool_ratios[2], gat_dims[1], 0.3)
502
+ self.pool_hT1 = GraphPool(pool_ratios[2], gat_dims[1], 0.3)
503
+
504
+ self.pool_hS2 = GraphPool(pool_ratios[2], gat_dims[1], 0.3)
505
+ self.pool_hT2 = GraphPool(pool_ratios[2], gat_dims[1], 0.3)
506
+
507
+ self.out_layer = nn.Linear(5 * gat_dims[1], 2)
508
+
509
+ def forward(self, x):
510
+ # -------pre-trained Wav2vec model fine tunning ------------------------##
511
+ x_ssl_feat = self.ssl_model.extract_feat(x.squeeze(-1))
512
+ x = self.LL(x_ssl_feat) # (bs,frame_number,feat_out_dim)
513
+
514
+ # post-processing on front-end features
515
+ x = x.transpose(1, 2) # (bs,feat_out_dim,frame_number)
516
+ x = x.unsqueeze(dim=1) # add channel
517
+ x = F.max_pool2d(x, (3, 3))
518
+ x = self.first_bn(x)
519
+ x = self.selu(x)
520
+
521
+ # RawNet2-based encoder
522
+ x = self.encoder(x)
523
+ x = self.first_bn1(x)
524
+ x = self.selu(x)
525
+
526
+ w = self.attention(x)
527
+
528
+ # ------------SA for spectral feature-------------#
529
+ w1 = F.softmax(w, dim=-1)
530
+ m = torch.sum(x * w1, dim=-1)
531
+ e_S = m.transpose(1, 2) + self.pos_S
532
+
533
+ # graph module layer
534
+ gat_S = self.GAT_layer_S(e_S)
535
+ out_S = self.pool_S(gat_S) # (#bs, #node, #dim)
536
+
537
+ # ------------SA for temporal feature-------------#
538
+ w2 = F.softmax(w, dim=-2)
539
+ m1 = torch.sum(x * w2, dim=-2)
540
+
541
+ e_T = m1.transpose(1, 2)
542
+
543
+ # graph module layer
544
+ gat_T = self.GAT_layer_T(e_T)
545
+ out_T = self.pool_T(gat_T)
546
+
547
+ # learnable master node
548
+ master1 = self.master1.expand(x.size(0), -1, -1)
549
+ master2 = self.master2.expand(x.size(0), -1, -1)
550
+
551
+ # inference 1
552
+ out_T1, out_S1, master1 = self.HtrgGAT_layer_ST11(
553
+ out_T, out_S, master=self.master1
554
+ )
555
+
556
+ out_S1 = self.pool_hS1(out_S1)
557
+ out_T1 = self.pool_hT1(out_T1)
558
+
559
+ out_T_aug, out_S_aug, master_aug = self.HtrgGAT_layer_ST12(
560
+ out_T1, out_S1, master=master1
561
+ )
562
+ out_T1 = out_T1 + out_T_aug
563
+ out_S1 = out_S1 + out_S_aug
564
+ master1 = master1 + master_aug
565
+
566
+ # inference 2
567
+ out_T2, out_S2, master2 = self.HtrgGAT_layer_ST21(
568
+ out_T, out_S, master=self.master2
569
+ )
570
+ out_S2 = self.pool_hS2(out_S2)
571
+ out_T2 = self.pool_hT2(out_T2)
572
+
573
+ out_T_aug, out_S_aug, master_aug = self.HtrgGAT_layer_ST22(
574
+ out_T2, out_S2, master=master2
575
+ )
576
+ out_T2 = out_T2 + out_T_aug
577
+ out_S2 = out_S2 + out_S_aug
578
+ master2 = master2 + master_aug
579
+
580
+ out_T1 = self.drop_way(out_T1)
581
+ out_T2 = self.drop_way(out_T2)
582
+ out_S1 = self.drop_way(out_S1)
583
+ out_S2 = self.drop_way(out_S2)
584
+ master1 = self.drop_way(master1)
585
+ master2 = self.drop_way(master2)
586
+
587
+ out_T = torch.max(out_T1, out_T2)
588
+ out_S = torch.max(out_S1, out_S2)
589
+ master = torch.max(master1, master2)
590
+
591
+ # Readout operation
592
+ T_max, _ = torch.max(torch.abs(out_T), dim=1)
593
+ T_avg = torch.mean(out_T, dim=1)
594
+
595
+ S_max, _ = torch.max(torch.abs(out_S), dim=1)
596
+ S_avg = torch.mean(out_S, dim=1)
597
+
598
+ last_hidden = torch.cat([T_max, T_avg, S_max, S_avg, master.squeeze(1)], dim=1)
599
+
600
+ last_hidden = self.drop(last_hidden)
601
+ output = self.out_layer(last_hidden)
602
+
603
+ return output
audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/requirements.txt ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ librosa==0.9.1
2
+ python-dotenv==1.0.1
3
+ tensorboardX
4
+ wandb
audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/startup_config.py ADDED
@@ -0,0 +1,60 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """
3
+ startup_config
4
+
5
+ Startup configuration utilities
6
+
7
+ """
8
+
9
+ from __future__ import absolute_import
10
+
11
+ import importlib
12
+ import os
13
+ import random
14
+ import sys
15
+
16
+ import numpy as np
17
+ import torch
18
+
19
+ __author__ = "Xin Wang"
20
+ __email__ = "wangxin@nii.ac.jp"
21
+ __copyright__ = "Copyright 2020, Xin Wang"
22
+
23
+
24
+ def set_random_seed(random_seed, args=None):
25
+ """set_random_seed(random_seed, args=None)
26
+
27
+ Set the random_seed for numpy, python, and cudnn
28
+
29
+ input
30
+ -----
31
+ random_seed: integer random seed
32
+ args: argue parser
33
+ """
34
+
35
+ # initialization
36
+ torch.manual_seed(random_seed)
37
+ random.seed(random_seed)
38
+ np.random.seed(random_seed)
39
+ os.environ["PYTHONHASHSEED"] = str(random_seed)
40
+
41
+ # For torch.backends.cudnn.deterministic
42
+ # Note: this default configuration may result in RuntimeError
43
+ # see https://pytorch.org/docs/stable/notes/randomness.html
44
+ if args is None:
45
+ cudnn_deterministic = True
46
+ cudnn_benchmark = False
47
+ else:
48
+ cudnn_deterministic = args.cudnn_deterministic_toggle
49
+ cudnn_benchmark = args.cudnn_benchmark_toggle
50
+
51
+ if not cudnn_deterministic:
52
+ print("cudnn_deterministic set to False")
53
+ if cudnn_benchmark:
54
+ print("cudnn_benchmark set to True")
55
+
56
+ if torch.cuda.is_available():
57
+ torch.cuda.manual_seed_all(random_seed)
58
+ torch.backends.cudnn.deterministic = cudnn_deterministic
59
+ torch.backends.cudnn.benchmark = cudnn_benchmark
60
+ return
audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/train.py ADDED
@@ -0,0 +1,446 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import os
3
+ import sys
4
+
5
+ import librosa
6
+ import numpy as np
7
+ import torch
8
+ import wandb
9
+ import yaml
10
+ from data_utils import (
11
+ Dataset_ASVspoof2019_train,
12
+ Dataset_ASVspoof2021_eval,
13
+ genSpoof_list,
14
+ pad,
15
+ process_Rawboost_feature,
16
+ )
17
+ from dotenv import load_dotenv
18
+ from model import Model
19
+ from sklearn.metrics import roc_auc_score
20
+ from startup_config import set_random_seed
21
+ from tensorboardX import SummaryWriter
22
+ from torch import Tensor, nn
23
+ from torch.utils.data import DataLoader
24
+ from tqdm import tqdm
25
+
26
+ __author__ = "Hemlata Tak"
27
+ __email__ = "tak@eurecom.fr"
28
+
29
+
30
+ def compute_det_curve(target_scores, nontarget_scores):
31
+
32
+ n_scores = target_scores.size + nontarget_scores.size
33
+ all_scores = np.concatenate((target_scores, nontarget_scores))
34
+ labels = np.concatenate(
35
+ (np.ones(target_scores.size), np.zeros(nontarget_scores.size))
36
+ )
37
+
38
+ indices = np.argsort(all_scores, kind="mergesort")
39
+ labels = labels[indices]
40
+ tar_trial_sums = np.cumsum(labels)
41
+ nontarget_trial_sums = nontarget_scores.size - (
42
+ np.arange(1, n_scores + 1) - tar_trial_sums
43
+ )
44
+
45
+ frr = np.concatenate((np.atleast_1d(0), tar_trial_sums / target_scores.size))
46
+ far = np.concatenate(
47
+ (np.atleast_1d(1), nontarget_trial_sums / nontarget_scores.size)
48
+ )
49
+ # Thresholds are the sorted scores
50
+ thresholds = np.concatenate(
51
+ (np.atleast_1d(all_scores[indices[0]] - 0.001), all_scores[indices])
52
+ )
53
+
54
+ return frr, far, thresholds
55
+
56
+
57
+ def compute_eer(target_scores, nontarget_scores):
58
+ """Returns equal error rate (EER) and the corresponding threshold."""
59
+ frr, far, thresholds = compute_det_curve(target_scores, nontarget_scores)
60
+ abs_diffs = np.abs(frr - far)
61
+ min_index = np.argmin(abs_diffs)
62
+ eer = np.mean((frr[min_index], far[min_index]))
63
+ return eer, thresholds[min_index], frr, far
64
+
65
+
66
+ def calculate_tDCF_EER(cm_scores_file, output_file, printout=True):
67
+ # Load CM scores
68
+ cm_data = np.genfromtxt(cm_scores_file, dtype=str)
69
+ cm_utt_id = cm_data[:, 0]
70
+ cm_keys = cm_data[:, 1]
71
+ cm_scores = cm_data[:, 2].astype(float)
72
+ # Extract bona fide (real human) and spoof scores from the CM scores
73
+ bona_cm = cm_scores[cm_keys == "bonafide"]
74
+ spoof_cm = cm_scores[cm_keys == "spoof"]
75
+ all_scores = np.concatenate([bona_cm, spoof_cm])
76
+ all_true_labels = np.concatenate([np.ones_like(bona_cm), np.zeros_like(spoof_cm)])
77
+
78
+ auc = roc_auc_score(all_true_labels, all_scores, max_fpr=0.05)
79
+ eer_cm, eer_threshold, frr, far = compute_eer(bona_cm, spoof_cm)
80
+
81
+ if printout:
82
+ with open(output_file, "w") as f_res:
83
+ f_res.write("\nCM SYSTEM\n")
84
+ f_res.write(
85
+ "\tEER\t\t= {:8.9f} % "
86
+ "(Equal error rate for countermeasure)\n".format(eer_cm * 100)
87
+ )
88
+ f_res.write("\t pAUC with max fpr - 0.05 is :{}".format(auc))
89
+
90
+
91
+ def evaluate_accuracy(dev_loader, model, device, args):
92
+ val_loss = 0.0
93
+ num_total = 0.0
94
+ algo = args.algo
95
+ cut = 64600
96
+ model.eval()
97
+
98
+ weight = torch.FloatTensor([0.1, 0.9]).to(device)
99
+ criterion = nn.CrossEntropyLoss(weight=weight)
100
+ progress_bar = tqdm(dev_loader, desc=f"Epoch {epoch+1}/{args.num_epochs}")
101
+ for current_step, (batch_pths, batch_y) in enumerate(progress_bar):
102
+ batch_x = batch_pths
103
+ batch_size = batch_x.size(0)
104
+ num_total += batch_size
105
+ batch_x = batch_x.to(device)
106
+ batch_y = batch_y.view(-1).type(torch.int64).to(device)
107
+ batch_out = model(batch_x)
108
+
109
+ batch_loss = criterion(batch_out, batch_y)
110
+ val_loss += batch_loss.item() * batch_size
111
+
112
+ val_loss /= num_total
113
+
114
+ return val_loss
115
+
116
+
117
+ def produce_evaluation_file(dataset, model, device, save_path, trial_path):
118
+ data_loader = DataLoader(dataset, batch_size=10, shuffle=False, drop_last=False)
119
+ num_correct = 0.0
120
+ num_total = 0.0
121
+ model.eval()
122
+ with open(trial_path, "r") as f_trl:
123
+ trial_lines = f_trl.readlines()
124
+
125
+ fname_list = []
126
+ score_list = []
127
+
128
+ for batch_x, utt_id in data_loader:
129
+
130
+ batch_size = batch_x.size(0)
131
+ batch_x = batch_x.to(device)
132
+
133
+ batch_out = model(batch_x)
134
+
135
+ batch_score = (batch_out[:, 1]).data.cpu().numpy().ravel()
136
+ # add outputs
137
+ fname_list.extend(utt_id)
138
+ score_list.extend(batch_score.tolist())
139
+ assert len(trial_lines) == len(fname_list) == len(score_list)
140
+
141
+ with open(save_path, "a+") as fh:
142
+ for fname, cm, trl in zip(fname_list, score_list, trial_lines):
143
+ utt_id, key = trl.strip().split(" ")
144
+ assert fname == utt_id
145
+ fh.write("{} {} {}\n".format(fname, key, cm))
146
+ fh.close()
147
+ print("Scores saved to {}".format(save_path))
148
+
149
+
150
+ def train_epoch(train_loader, model, lr, optim, device, args):
151
+ running_loss = 0
152
+
153
+ num_total = 0.0
154
+ algo = args.algo
155
+ model.train()
156
+ cut = 64600
157
+ # set objective (Loss) functions
158
+ weight = torch.FloatTensor([0.1, 0.9]).to(device)
159
+ criterion = nn.CrossEntropyLoss(weight=weight)
160
+ progress_bar = tqdm(train_loader, desc=f"Epoch {epoch+1}/{args.num_epochs}")
161
+ for current_step, (batch_pths, batch_y) in enumerate(progress_bar):
162
+ batch_x = batch_pths
163
+ batch_size = batch_x.size(0)
164
+ num_total += batch_size
165
+
166
+ batch_x = batch_x.to(device)
167
+ batch_y = batch_y.view(-1).type(torch.int64).to(device)
168
+ batch_out = model(batch_x)
169
+
170
+ batch_loss = criterion(batch_out, batch_y)
171
+
172
+ running_loss += batch_loss.item() * batch_size
173
+
174
+ optimizer.zero_grad()
175
+ batch_loss.backward()
176
+ optimizer.step()
177
+
178
+ running_loss /= num_total
179
+
180
+ return running_loss
181
+
182
+
183
+ if __name__ == "__main__":
184
+ parser = argparse.ArgumentParser(description="SSL-AASIST baseline system")
185
+
186
+ # Hyperparameters
187
+ parser.add_argument("--batch_size", type=int, default=64)
188
+ parser.add_argument("--num_epochs", type=int, default=100)
189
+ parser.add_argument("--lr", type=float, default=0.000001)
190
+ parser.add_argument("--weight_decay", type=float, default=0.0001)
191
+ parser.add_argument("--model_name", type=str, default="SSL-AASIST")
192
+ parser.add_argument("--loss", type=str, default="weighted_CCE")
193
+ parser.add_argument("--trn_list_path", default=None, help="path to train file")
194
+ parser.add_argument("--dev_list_path", default=None, help="path to validation file")
195
+ parser.add_argument("--test_list_path", default=None, help="path to test file")
196
+ parser.add_argument(
197
+ "--test_score_dir", default=None, help="path to save test scores"
198
+ )
199
+ # model
200
+ parser.add_argument(
201
+ "--seed", type=int, default=1234, help="random seed (default: 1234)"
202
+ )
203
+ parser.add_argument("--save_path", type=str, default=".", help="Model save path")
204
+ parser.add_argument("--model_path", type=str, default=None, help="Model checkpoint")
205
+ parser.add_argument(
206
+ "--comment", type=str, default=None, help="Comment to describe the saved model"
207
+ )
208
+ # Auxiliary arguments
209
+
210
+ parser.add_argument("--eval", action="store_true", default=False, help="eval mode")
211
+ parser.add_argument("--eval_part", type=int, default=0)
212
+ # backend options
213
+ parser.add_argument(
214
+ "--cudnn-deterministic-toggle",
215
+ action="store_false",
216
+ default=True,
217
+ help="use cudnn-deterministic? (default true)",
218
+ )
219
+
220
+ parser.add_argument(
221
+ "--cudnn-benchmark-toggle",
222
+ action="store_true",
223
+ default=False,
224
+ help="use cudnn-benchmark? (default false)",
225
+ )
226
+
227
+ ##===================================================Rawboost data augmentation ======================================================================#
228
+
229
+ parser.add_argument(
230
+ "--algo",
231
+ type=int,
232
+ default=5,
233
+ help="Rawboost algos discriptions. 0: No augmentation 1: LnL_convolutive_noise, 2: ISD_additive_noise, 3: SSI_additive_noise, 4: series algo (1+2+3), \
234
+ 5: series algo (1+2), 6: series algo (1+3), 7: series algo(2+3), 8: parallel algo(1,2) .[default=0]",
235
+ )
236
+
237
+ # LnL_convolutive_noise parameters
238
+ parser.add_argument(
239
+ "--nBands",
240
+ type=int,
241
+ default=5,
242
+ help="number of notch filters.The higher the number of bands, the more aggresive the distortions is.[default=5]",
243
+ )
244
+ parser.add_argument(
245
+ "--minF",
246
+ type=int,
247
+ default=20,
248
+ help="minimum centre frequency [Hz] of notch filter.[default=20] ",
249
+ )
250
+ parser.add_argument(
251
+ "--maxF",
252
+ type=int,
253
+ default=8000,
254
+ help="maximum centre frequency [Hz] (<sr/2) of notch filter.[default=8000]",
255
+ )
256
+ parser.add_argument(
257
+ "--minBW",
258
+ type=int,
259
+ default=100,
260
+ help="minimum width [Hz] of filter.[default=100] ",
261
+ )
262
+ parser.add_argument(
263
+ "--maxBW",
264
+ type=int,
265
+ default=1000,
266
+ help="maximum width [Hz] of filter.[default=1000] ",
267
+ )
268
+ parser.add_argument(
269
+ "--minCoeff",
270
+ type=int,
271
+ default=10,
272
+ help="minimum filter coefficients. More the filter coefficients more ideal the filter slope.[default=10]",
273
+ )
274
+ parser.add_argument(
275
+ "--maxCoeff",
276
+ type=int,
277
+ default=100,
278
+ help="maximum filter coefficients. More the filter coefficients more ideal the filter slope.[default=100]",
279
+ )
280
+ parser.add_argument(
281
+ "--minG",
282
+ type=int,
283
+ default=0,
284
+ help="minimum gain factor of linear component.[default=0]",
285
+ )
286
+ parser.add_argument(
287
+ "--maxG",
288
+ type=int,
289
+ default=0,
290
+ help="maximum gain factor of linear component.[default=0]",
291
+ )
292
+ parser.add_argument(
293
+ "--minBiasLinNonLin",
294
+ type=int,
295
+ default=5,
296
+ help=" minimum gain difference between linear and non-linear components.[default=5]",
297
+ )
298
+ parser.add_argument(
299
+ "--maxBiasLinNonLin",
300
+ type=int,
301
+ default=20,
302
+ help=" maximum gain difference between linear and non-linear components.[default=20]",
303
+ )
304
+ parser.add_argument(
305
+ "--N_f",
306
+ type=int,
307
+ default=5,
308
+ help="order of the (non-)linearity where N_f=1 refers only to linear components.[default=5]",
309
+ )
310
+
311
+ # ISD_additive_noise parameters
312
+ parser.add_argument(
313
+ "--P",
314
+ type=int,
315
+ default=10,
316
+ help="Maximum number of uniformly distributed samples in [%].[defaul=10]",
317
+ )
318
+ parser.add_argument(
319
+ "--g_sd", type=int, default=2, help="gain parameters > 0. [default=2]"
320
+ )
321
+
322
+ # SSI_additive_noise parameters
323
+ parser.add_argument(
324
+ "--SNRmin",
325
+ type=int,
326
+ default=10,
327
+ help="Minimum SNR value for coloured additive noise.[defaul=10]",
328
+ )
329
+ parser.add_argument(
330
+ "--SNRmax",
331
+ type=int,
332
+ default=40,
333
+ help="Maximum SNR value for coloured additive noise.[defaul=40]",
334
+ )
335
+
336
+ ##===================================================Rawboost data augmentation ======================================================================#
337
+
338
+ load_dotenv()
339
+ wandb_api_key = os.getenv("WANDB_API_KEY")
340
+ wandb_project_name = os.getenv("WANDB_PROJECT_NAME")
341
+
342
+ if not os.path.exists("models"):
343
+ os.mkdir("models")
344
+ args = parser.parse_args()
345
+ wandb.login(key=wandb_api_key)
346
+ wandb.init(
347
+ project=wandb_project_name,
348
+ config={
349
+ "learning_rate": args.lr,
350
+ "epochs": args.num_epochs,
351
+ "batch_size": args.batch_size,
352
+ "weight_decay": args.weight_decay,
353
+ },
354
+ )
355
+
356
+ # make experiment reproducible
357
+ set_random_seed(args.seed, args)
358
+
359
+ # define model saving path
360
+ model_tag = "model_{}_{}_{}_{}".format(
361
+ args.loss, args.num_epochs, args.batch_size, args.lr
362
+ )
363
+ if args.comment:
364
+ model_tag = model_tag + "_{}".format(args.comment)
365
+ model_save_path = os.path.join(args.save_path, model_tag)
366
+
367
+ # set model save directory
368
+ if not os.path.exists(model_save_path):
369
+ os.mkdir(model_save_path)
370
+
371
+ # GPU device
372
+ device = "cuda" if torch.cuda.is_available() else "cpu"
373
+ print("Device: {}".format(device))
374
+
375
+ model = Model(args, device)
376
+ nb_params = sum([param.view(-1).size()[0] for param in model.parameters()])
377
+ model = model.to(device)
378
+ print("nb_params:", nb_params)
379
+
380
+ # set Adam optimizer
381
+ optimizer = torch.optim.Adam(
382
+ model.parameters(), lr=args.lr, weight_decay=args.weight_decay
383
+ )
384
+
385
+ if args.model_path:
386
+ model.load_state_dict(torch.load(args.model_path, map_location=device))
387
+ print("Model loaded : {}".format(args.model_path))
388
+
389
+ # evaluation
390
+
391
+ if args.eval:
392
+ file_eval = genSpoof_list(
393
+ dir_meta=args.test_list_path, is_train=False, is_eval=True
394
+ )
395
+ print("no. of eval trials", len(file_eval))
396
+ eval_set = Dataset_ASVspoof2021_eval(list_IDs=file_eval)
397
+ eval_output = os.path.join(
398
+ args.test_score_dir, f"{args.model_name}_model_score.txt"
399
+ )
400
+ produce_evaluation_file(
401
+ eval_set, model, device, eval_output, args.test_list_path
402
+ )
403
+ output_file = os.path.join(
404
+ args.test_score_dir, f"{args.model_name}_model_eer.txt"
405
+ )
406
+ eval_eer = calculate_tDCF_EER(
407
+ cm_scores_file=eval_output, output_file=output_file
408
+ )
409
+
410
+ sys.exit(0)
411
+
412
+ trn_list_path = args.trn_list_path
413
+ dev_trial_path = args.dev_list_path
414
+ train_set = Dataset_ASVspoof2019_train(args, metafile=trn_list_path, algo=args.algo)
415
+ train_loader = DataLoader(
416
+ train_set,
417
+ batch_size=args.batch_size,
418
+ num_workers=16,
419
+ shuffle=True,
420
+ drop_last=True,
421
+ )
422
+ del train_set
423
+
424
+ dev_set = Dataset_ASVspoof2019_train(args, metafile=dev_trial_path, algo=args.algo)
425
+ dev_loader = DataLoader(
426
+ dev_set, batch_size=args.batch_size, num_workers=16, shuffle=False
427
+ )
428
+ del dev_set
429
+ # Training and validation
430
+ num_epochs = args.num_epochs
431
+ writer = SummaryWriter("logs/{}".format(model_tag))
432
+
433
+ for epoch in range(num_epochs):
434
+
435
+ running_loss = train_epoch(
436
+ train_loader, model, args.lr, optimizer, device, args
437
+ )
438
+ val_loss = evaluate_accuracy(dev_loader, model, device, args)
439
+ wandb.log({"epoch": epoch, "train_loss": running_loss, "val_loss": val_loss})
440
+ writer.add_scalar("val_loss", val_loss, epoch)
441
+ writer.add_scalar("loss", running_loss, epoch)
442
+ print("\n{} - {} - {} ".format(epoch, running_loss, val_loss))
443
+ torch.save(
444
+ model.state_dict(),
445
+ os.path.join(model_save_path, "epoch_{}.pth".format(epoch)),
446
+ )
audio/shiftyspeech/tests/__init__.py ADDED
File without changes
audio/shiftyspeech/tests/test_api.py ADDED
@@ -0,0 +1,253 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tests for ShiftySpeech Audio Deepfake Detection API.
2
+
3
+ Tests cover:
4
+ - Health endpoint
5
+ - Predict endpoint with real audio
6
+ - Predict endpoint with fake audio
7
+ - Audio preprocessing (pad/trim/tile)
8
+ - Error handling for invalid input
9
+ - Response schema validation
10
+ """
11
+
12
+ import base64
13
+ import io
14
+ import os
15
+ import sys
16
+ import warnings
17
+
18
+ import numpy as np
19
+ import pytest
20
+
21
+ warnings.filterwarnings("ignore", category=DeprecationWarning)
22
+
23
+ # Monkey-patch omegaconf before importing api module
24
+ import omegaconf._utils as _omegaconf_utils
25
+
26
+ if not hasattr(_omegaconf_utils, "is_primitive_type"):
27
+ _omegaconf_utils.is_primitive_type = lambda t: t in (int, float, bool, str, bytes)
28
+
29
+ # Add model code path for local testing
30
+ SERVICE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
31
+ MODEL_CODE_PATH = os.path.join(
32
+ SERVICE_DIR, "synthetic_speech_detection", "SSL_Anti-spoofing"
33
+ )
34
+ if MODEL_CODE_PATH not in sys.path:
35
+ sys.path.insert(0, MODEL_CODE_PATH)
36
+
37
+ # Patch the api module constants for local testing
38
+ import api
39
+
40
+ api.MODEL_CODE_PATH = MODEL_CODE_PATH
41
+ api.WEIGHTS_PATH = os.path.join(SERVICE_DIR, "weights", "hfg_aug_1_2.pt")
42
+ api.XLSR_DIR = os.path.join(SERVICE_DIR, "models")
43
+
44
+ from fastapi.testclient import TestClient
45
+
46
+ client = TestClient(api.app)
47
+
48
+ # Dataset paths
49
+ DATASET_DIR = os.path.join(
50
+ os.path.dirname(SERVICE_DIR),
51
+ os.pardir,
52
+ os.pardir,
53
+ "dataset",
54
+ "audio",
55
+ )
56
+ DATASET_DIR = os.path.normpath(DATASET_DIR)
57
+ REAL_DIR = os.path.join(DATASET_DIR, "real")
58
+ FAKE_DIR = os.path.join(DATASET_DIR, "fake")
59
+
60
+ # Check if model weights are available for integration tests
61
+ WEIGHTS_AVAILABLE = os.path.exists(api.WEIGHTS_PATH) and os.path.exists(
62
+ os.path.join(SERVICE_DIR, "models", "xlsr2_300m.pt")
63
+ )
64
+ DATASET_AVAILABLE = os.path.isdir(REAL_DIR) and os.path.isdir(FAKE_DIR)
65
+
66
+
67
+ def _make_wav_bytes(duration_s: float = 1.0, sr: int = 16000) -> bytes:
68
+ """Generate a simple sine wave WAV file as bytes."""
69
+ import soundfile as sf
70
+
71
+ t = np.linspace(0, duration_s, int(sr * duration_s), endpoint=False)
72
+ audio = 0.5 * np.sin(2 * np.pi * 440 * t).astype(np.float32)
73
+ buf = io.BytesIO()
74
+ sf.write(buf, audio, sr, format="WAV")
75
+ buf.seek(0)
76
+ return buf.read()
77
+
78
+
79
+ def _encode_file(path: str) -> str:
80
+ """Read a file and return base64 encoded string."""
81
+ with open(path, "rb") as f:
82
+ return base64.b64encode(f.read()).decode("utf-8")
83
+
84
+
85
+ class TestHealthEndpoint:
86
+ """Tests for the /health endpoint."""
87
+
88
+ def test_health_returns_200(self):
89
+ response = client.get("/health")
90
+ assert response.status_code == 200
91
+
92
+ def test_health_contains_model_name(self):
93
+ response = client.get("/health")
94
+ data = response.json()
95
+ assert data["model"] == "shiftyspeech"
96
+
97
+ def test_health_contains_device(self):
98
+ response = client.get("/health")
99
+ data = response.json()
100
+ assert data["device"] == "cpu"
101
+
102
+ def test_health_contains_status(self):
103
+ response = client.get("/health")
104
+ data = response.json()
105
+ assert data["status"] in ("healthy", "degraded")
106
+
107
+
108
+ class TestPreprocessAudio:
109
+ """Tests for audio preprocessing logic."""
110
+
111
+ def test_preprocess_short_audio_tiles(self):
112
+ """Short audio should be tiled to TARGET_SAMPLES."""
113
+ wav_bytes = _make_wav_bytes(duration_s=0.5, sr=16000)
114
+ tensor = api.preprocess_audio(wav_bytes)
115
+ assert tensor.shape == (1, api.TARGET_SAMPLES)
116
+
117
+ def test_preprocess_long_audio_trims(self):
118
+ """Long audio should be trimmed to TARGET_SAMPLES."""
119
+ wav_bytes = _make_wav_bytes(duration_s=10.0, sr=16000)
120
+ tensor = api.preprocess_audio(wav_bytes)
121
+ assert tensor.shape == (1, api.TARGET_SAMPLES)
122
+
123
+ def test_preprocess_exact_length(self):
124
+ """Audio at exact TARGET_SAMPLES should pass through."""
125
+ duration = api.TARGET_SAMPLES / api.SAMPLE_RATE
126
+ wav_bytes = _make_wav_bytes(duration_s=duration, sr=16000)
127
+ tensor = api.preprocess_audio(wav_bytes)
128
+ assert tensor.shape == (1, api.TARGET_SAMPLES)
129
+
130
+ def test_preprocess_resamples_from_8khz(self):
131
+ """Audio at 8kHz should be resampled to 16kHz."""
132
+ wav_bytes = _make_wav_bytes(duration_s=1.0, sr=8000)
133
+ tensor = api.preprocess_audio(wav_bytes)
134
+ assert tensor.shape == (1, api.TARGET_SAMPLES)
135
+
136
+ def test_preprocess_invalid_input_raises(self):
137
+ """Invalid audio bytes should raise ValueError."""
138
+ with pytest.raises(ValueError):
139
+ api.preprocess_audio(b"not audio data")
140
+
141
+
142
+ @pytest.mark.skipif(
143
+ not WEIGHTS_AVAILABLE,
144
+ reason="Model weights not available locally",
145
+ )
146
+ class TestPredictEndpoint:
147
+ """Integration tests for the /predict endpoint (requires weights)."""
148
+
149
+ def test_predict_returns_200(self):
150
+ wav_bytes = _make_wav_bytes(duration_s=2.0)
151
+ b64 = base64.b64encode(wav_bytes).decode("utf-8")
152
+ response = client.post("/predict", json={"audio_data": b64})
153
+ assert response.status_code == 200
154
+
155
+ def test_predict_response_schema(self):
156
+ wav_bytes = _make_wav_bytes(duration_s=2.0)
157
+ b64 = base64.b64encode(wav_bytes).decode("utf-8")
158
+ response = client.post("/predict", json={"audio_data": b64})
159
+ data = response.json()
160
+ assert "model" in data
161
+ assert "probability" in data
162
+ assert "prediction" in data
163
+ assert "class" in data
164
+ assert "inference_time" in data
165
+ assert data["model"] == "shiftyspeech"
166
+
167
+ def test_predict_probability_in_range(self):
168
+ wav_bytes = _make_wav_bytes(duration_s=2.0)
169
+ b64 = base64.b64encode(wav_bytes).decode("utf-8")
170
+ response = client.post("/predict", json={"audio_data": b64})
171
+ data = response.json()
172
+ assert 0.0 <= data["probability"] <= 1.0
173
+
174
+ def test_predict_class_matches_prediction(self):
175
+ wav_bytes = _make_wav_bytes(duration_s=2.0)
176
+ b64 = base64.b64encode(wav_bytes).decode("utf-8")
177
+ response = client.post("/predict", json={"audio_data": b64})
178
+ data = response.json()
179
+ if data["prediction"] == 1:
180
+ assert data["class"] == "fake"
181
+ else:
182
+ assert data["class"] == "real"
183
+
184
+ def test_predict_custom_threshold(self):
185
+ wav_bytes = _make_wav_bytes(duration_s=2.0)
186
+ b64 = base64.b64encode(wav_bytes).decode("utf-8")
187
+ response = client.post(
188
+ "/predict",
189
+ json={"audio_data": b64, "threshold": 0.99},
190
+ )
191
+ data = response.json()
192
+ assert response.status_code == 200
193
+ # With threshold=0.99, only very high prob_fake => fake
194
+ if data["probability"] < 0.99:
195
+ assert data["prediction"] == 0
196
+ assert data["class"] == "real"
197
+
198
+ def test_predict_inference_time_positive(self):
199
+ wav_bytes = _make_wav_bytes(duration_s=2.0)
200
+ b64 = base64.b64encode(wav_bytes).decode("utf-8")
201
+ response = client.post("/predict", json={"audio_data": b64})
202
+ data = response.json()
203
+ assert data["inference_time"] > 0
204
+
205
+
206
+ @pytest.mark.skipif(
207
+ not WEIGHTS_AVAILABLE or not DATASET_AVAILABLE,
208
+ reason="Model weights or dataset not available",
209
+ )
210
+ class TestRealDataset:
211
+ """Integration tests using actual dataset files."""
212
+
213
+ def test_predict_real_audio(self):
214
+ """Test prediction on a real audio file."""
215
+ path = os.path.join(REAL_DIR, "real_0.wav")
216
+ if not os.path.exists(path):
217
+ pytest.skip("real_0.wav not found")
218
+ b64 = _encode_file(path)
219
+ response = client.post("/predict", json={"audio_data": b64})
220
+ assert response.status_code == 200
221
+ data = response.json()
222
+ assert 0.0 <= data["probability"] <= 1.0
223
+
224
+ def test_predict_fake_audio(self):
225
+ """Test prediction on a fake audio file."""
226
+ path = os.path.join(FAKE_DIR, "fake_1.wav")
227
+ if not os.path.exists(path):
228
+ pytest.skip("fake_1.wav not found")
229
+ b64 = _encode_file(path)
230
+ response = client.post("/predict", json={"audio_data": b64})
231
+ assert response.status_code == 200
232
+ data = response.json()
233
+ assert 0.0 <= data["probability"] <= 1.0
234
+
235
+
236
+ class TestPredictValidation:
237
+ """Tests for input validation on /predict."""
238
+
239
+ def test_predict_missing_audio_data(self):
240
+ response = client.post("/predict", json={})
241
+ assert response.status_code == 422
242
+
243
+ def test_predict_invalid_base64(self):
244
+ response = client.post("/predict", json={"audio_data": "not-valid-base64!!!"})
245
+ # Should return 500 (decode error) or 422
246
+ assert response.status_code in (400, 422, 500)
247
+
248
+ def test_predict_threshold_out_of_range(self):
249
+ response = client.post(
250
+ "/predict",
251
+ json={"audio_data": "dGVzdA==", "threshold": 1.5},
252
+ )
253
+ assert response.status_code == 422
audio/sonics/Dockerfile ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM nvidia/cuda:12.1.1-cudnn8-runtime-ubuntu22.04
2
+
3
+ ENV DEBIAN_FRONTEND=noninteractive
4
+ ENV PYTHONUNBUFFERED=1
5
+
6
+ RUN apt-get update && apt-get install -y --no-install-recommends \
7
+ python3 python3-pip python3-dev \
8
+ git ffmpeg libsndfile1 \
9
+ build-essential g++ \
10
+ && rm -rf /var/lib/apt/lists/*
11
+
12
+ RUN ln -sf /usr/bin/python3 /usr/bin/python
13
+
14
+ WORKDIR /app
15
+
16
+ # Install PyTorch with CUDA 12.1
17
+ RUN pip install --no-cache-dir \
18
+ torch==2.5.1 torchaudio==2.5.1 \
19
+ --index-url https://download.pytorch.org/whl/cu121
20
+
21
+ COPY requirements.txt .
22
+ RUN pip install --no-cache-dir -r requirements.txt
23
+
24
+ COPY app.py .
25
+
26
+ # Pre-download model weights
27
+ RUN python -c "from sonics import HFAudioClassifier; HFAudioClassifier.from_pretrained('awsaf49/sonics-spectttra-alpha-120s')"
28
+
29
+ EXPOSE 8003
30
+
31
+ RUN adduser --disabled-password --gecos '' appuser
32
+ USER appuser
33
+
34
+ ENV PRELOAD_MODEL=true
35
+
36
+ CMD ["python", "app.py"]
audio/sonics/app.py ADDED
@@ -0,0 +1,285 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """SONICS (SpecTTTra) Synthetic Music Detection API.
2
+
3
+ Detects AI-generated music (Suno, Udio, etc.) using the SpecTTTra
4
+ architecture from the SONICS project (ICLR 2025).
5
+
6
+ The model performs binary classification on raw audio waveforms via
7
+ internal MelSpectrogram features. It outputs a single logit; we apply
8
+ sigmoid to obtain the fake probability.
9
+
10
+ Reference: https://github.com/awsaf49/sonics
11
+ """
12
+
13
+ import base64
14
+ import io
15
+ import logging
16
+ import os
17
+ import platform
18
+ import time
19
+ from typing import Optional
20
+
21
+ import librosa
22
+ import numpy as np
23
+ import torch
24
+ import uvicorn
25
+ from fastapi import FastAPI, HTTPException
26
+ from pydantic import BaseModel, Field
27
+
28
+ # Configure logging
29
+ logging.basicConfig(
30
+ level=logging.INFO,
31
+ format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
32
+ )
33
+ logger = logging.getLogger("sonics_api")
34
+
35
+ # Constants
36
+ MODEL_NAME = "sonics_detection"
37
+ HF_MODEL_ID = "awsaf49/sonics-spectttra-alpha-120s"
38
+ SAMPLE_RATE = 16000
39
+ MAX_TIME = 120 # seconds (matches alpha-120s config)
40
+ MAX_LEN = MAX_TIME * SAMPLE_RATE # 1_920_000 samples
41
+ PRELOAD_MODEL = os.environ.get("PRELOAD_MODEL", "true").lower() == "true"
42
+ MODEL_TIMEOUT = int(os.environ.get("MODEL_TIMEOUT", "300"))
43
+
44
+
45
+ def _get_device():
46
+ """Select optimal device: MPS (Apple) > CUDA (NVIDIA) > CPU."""
47
+ override = os.environ.get("DEEPSAFE_DEVICE", "").strip().lower()
48
+ if override == "cpu":
49
+ return torch.device("cpu")
50
+ if override == "cuda" and torch.cuda.is_available():
51
+ return torch.device("cuda")
52
+ if (
53
+ override == "mps"
54
+ and hasattr(torch.backends, "mps")
55
+ and torch.backends.mps.is_available()
56
+ ):
57
+ return torch.device("mps")
58
+ if override:
59
+ pass # Invalid override, fall through to auto-detect
60
+ if (
61
+ platform.system() == "Darwin"
62
+ and hasattr(torch.backends, "mps")
63
+ and torch.backends.mps.is_available()
64
+ ):
65
+ return torch.device("mps")
66
+ if torch.cuda.is_available():
67
+ return torch.device("cuda")
68
+ return torch.device("cpu")
69
+
70
+
71
+ DEVICE = _get_device()
72
+
73
+ if DEVICE.type == "cuda":
74
+ torch.backends.cudnn.benchmark = True
75
+ torch.set_float32_matmul_precision("high")
76
+
77
+ if DEVICE.type == "cuda":
78
+ logger.info(
79
+ "Device: cuda (%s, %.1f GB VRAM)",
80
+ torch.cuda.get_device_name(0),
81
+ torch.cuda.get_device_properties(0).total_memory / 1024**3,
82
+ )
83
+ else:
84
+ logger.warning(
85
+ "Device: %s (no CUDA available -- check nvidia-container-toolkit)",
86
+ DEVICE,
87
+ )
88
+
89
+ # Global model instance
90
+ model = None
91
+
92
+
93
+ class AudioInput(BaseModel):
94
+ """Request schema for audio deepfake detection."""
95
+
96
+ audio_data: str = Field(
97
+ ..., description="Base64 encoded audio string (WAV/MP3/FLAC/etc)"
98
+ )
99
+ threshold: Optional[float] = Field(
100
+ 0.5, ge=0.0, le=1.0, description="Classification threshold"
101
+ )
102
+
103
+
104
+ app = FastAPI(
105
+ title="SONICS Synthetic Music Detection API",
106
+ description=(
107
+ "Service for detecting AI-generated music using the "
108
+ "SpecTTTra model from the SONICS project (ICLR 2025)."
109
+ ),
110
+ version="1.0.0",
111
+ )
112
+
113
+
114
+ def load_model():
115
+ """Load the SONICS HFAudioClassifier from HuggingFace Hub.
116
+
117
+ Returns:
118
+ The loaded model, or None if loading fails.
119
+ """
120
+ global model
121
+ if model is not None:
122
+ return model
123
+
124
+ logger.info("Loading SONICS model '%s' onto %s...", HF_MODEL_ID, DEVICE)
125
+
126
+ try:
127
+ from sonics import HFAudioClassifier
128
+
129
+ model = HFAudioClassifier.from_pretrained(
130
+ HF_MODEL_ID,
131
+ map_location=str(DEVICE),
132
+ )
133
+ model.to(DEVICE)
134
+ model.eval()
135
+
136
+ logger.info("SONICS model loaded successfully.")
137
+ return model
138
+ except Exception:
139
+ logger.exception("Failed to load SONICS model")
140
+ model = None
141
+ return None
142
+
143
+
144
+ @app.on_event("startup")
145
+ async def startup_event():
146
+ """Optionally preload model on service startup."""
147
+ if PRELOAD_MODEL:
148
+ load_model()
149
+
150
+
151
+ @app.get("/")
152
+ async def root():
153
+ """Root info endpoint."""
154
+ return {
155
+ "service": "SONICS Synthetic Music Detection",
156
+ "model": MODEL_NAME,
157
+ "version": "1.0.0",
158
+ }
159
+
160
+
161
+ def _gpu_health_info() -> dict:
162
+ """Return GPU metrics for the health endpoint."""
163
+ if torch.cuda.is_available() and DEVICE.type == "cuda":
164
+ return {
165
+ "gpu_name": torch.cuda.get_device_name(0),
166
+ "vram_used_mb": round(torch.cuda.memory_allocated(0) / 1024**2),
167
+ "vram_total_mb": round(
168
+ torch.cuda.get_device_properties(0).total_memory / 1024**2
169
+ ),
170
+ }
171
+ return {}
172
+
173
+
174
+ @app.get("/health")
175
+ async def health():
176
+ """Health check endpoint."""
177
+ return {
178
+ "status": "healthy" if model is not None else "degraded",
179
+ "model": MODEL_NAME,
180
+ "device": str(DEVICE),
181
+ **_gpu_health_info(),
182
+ }
183
+
184
+
185
+ def preprocess_audio(audio_bytes: bytes) -> torch.Tensor:
186
+ """Preprocess audio for SONICS inference.
187
+
188
+ Loads audio from raw bytes, resamples to 16 kHz mono,
189
+ crops or zero-pads to MAX_LEN samples, and normalises by
190
+ standard deviation (matching the training pipeline).
191
+
192
+ Args:
193
+ audio_bytes: Raw audio file bytes (WAV, MP3, FLAC, etc.).
194
+
195
+ Returns:
196
+ Audio tensor of shape (1, MAX_LEN) on DEVICE.
197
+
198
+ Raises:
199
+ ValueError: If audio preprocessing fails.
200
+ """
201
+ try:
202
+ logger.info("Starting audio preprocessing...")
203
+ audio, sr = librosa.load(io.BytesIO(audio_bytes), sr=SAMPLE_RATE, mono=True)
204
+ logger.info("Audio loaded. Length: %d samples at %dHz", len(audio), sr)
205
+
206
+ # Crop or pad to fixed length (matching SONICS dataset.py)
207
+ if len(audio) > MAX_LEN:
208
+ # Crop from 3/4 position (matching eval-mode logic)
209
+ idx = int((len(audio) - MAX_LEN) / 4 * 3)
210
+ audio = audio[idx : idx + MAX_LEN]
211
+ elif len(audio) < MAX_LEN:
212
+ audio = np.pad(audio, (0, MAX_LEN - len(audio)), mode="constant")
213
+
214
+ # Normalise by standard deviation (matching training pipeline)
215
+ audio /= np.maximum(np.std(audio), 1e-6)
216
+
217
+ logger.info("Audio preprocessed to %d samples", len(audio))
218
+
219
+ audio_tensor = torch.from_numpy(audio).float().unsqueeze(0)
220
+ audio_tensor = audio_tensor.to(DEVICE)
221
+ return audio_tensor
222
+ except Exception as e:
223
+ logger.error("Error preprocessing audio: %s", e)
224
+ raise ValueError(f"Audio preprocessing failed: {str(e)}")
225
+
226
+
227
+ @app.post("/predict")
228
+ async def predict(input_data: AudioInput):
229
+ """Run synthetic music detection on base64-encoded audio.
230
+
231
+ The model uses BCEWithLogitsLoss with num_classes=1, so it
232
+ outputs a single logit. We apply sigmoid to obtain the fake
233
+ probability.
234
+ """
235
+ if model is None:
236
+ if load_model() is None:
237
+ raise HTTPException(status_code=503, detail="Model not loaded")
238
+
239
+ try:
240
+ start_time = time.time()
241
+ logger.info(
242
+ "Prediction request. Data size: %d chars",
243
+ len(input_data.audio_data),
244
+ )
245
+
246
+ # Decode base64 audio
247
+ audio_bytes = base64.b64decode(input_data.audio_data)
248
+
249
+ # Preprocess
250
+ audio_tensor = preprocess_audio(audio_bytes)
251
+
252
+ # Inference
253
+ logger.info("Starting model inference...")
254
+ with torch.no_grad():
255
+ logits = model(audio_tensor)
256
+ # logits shape: (1, 1) -- single logit for binary classification
257
+ prob_fake = torch.sigmoid(logits).squeeze().item()
258
+
259
+ prediction = 1 if prob_fake >= input_data.threshold else 0
260
+ verdict = "fake" if prediction == 1 else "real"
261
+ inference_time = time.time() - start_time
262
+
263
+ logger.info(
264
+ "Prediction: %s (prob_fake=%.4f, time=%.3fs)",
265
+ verdict,
266
+ prob_fake,
267
+ inference_time,
268
+ )
269
+
270
+ return {
271
+ "model": MODEL_NAME,
272
+ "probability": float(prob_fake),
273
+ "prediction": int(prediction),
274
+ "class": verdict,
275
+ "inference_time": float(inference_time),
276
+ }
277
+
278
+ except Exception as e:
279
+ logger.exception("Error during prediction: %s", e)
280
+ raise HTTPException(status_code=500, detail=str(e))
281
+
282
+
283
+ if __name__ == "__main__":
284
+ port = int(os.environ.get("MODEL_PORT", 8003))
285
+ uvicorn.run(app, host="0.0.0.0", port=port)
audio/sonics/requirements.txt ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SONICS model (pip-installable from GitHub)
2
+ sonics @ git+https://github.com/awsaf49/sonics.git
3
+
4
+ # Core inference
5
+ torch>=2.4.0
6
+ torchaudio>=2.4.0
7
+
8
+ # Audio processing
9
+ librosa>=0.9.0
10
+
11
+ # ML utilities (SONICS dependency)
12
+ timm>=1.0.7
13
+ fvcore
14
+
15
+ # HuggingFace Hub (for model download)
16
+ huggingface-hub
17
+
18
+ # API
19
+ fastapi
20
+ uvicorn[standard]
21
+ pydantic>=2.0
22
+ numpy
ensemble-core/Dockerfile ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM python:3.9-slim
2
+ WORKDIR /app
3
+ COPY requirements.txt .
4
+ RUN pip install --no-cache-dir -r requirements.txt
5
+ COPY main.py .
6
+ # Drop root privileges
7
+ RUN adduser --disabled-password --gecos '' appuser
8
+ USER appuser
9
+
10
+ CMD ["uvicorn", "main:app", "--host", "0.0.0.0", "--port", "8003"]
ensemble-core/main.py ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Any, Dict
2
+
3
+ from fastapi import FastAPI
4
+ from pydantic import BaseModel
5
+
6
+ app = FastAPI(title="Ensemble Core Service")
7
+
8
+
9
+ class EnsembleRequest(BaseModel):
10
+ media_type: str
11
+ model_results: Dict[str, Any]
12
+ method: str = "stacking"
13
+
14
+
15
+ @app.post("/calculate")
16
+ async def calculate_ensemble(request: EnsembleRequest):
17
+ # TODO: Migrate proprietary ensemble logic from gateway
18
+ return {"verdict": "fake", "confidence": 0.95, "method_used": request.method}
19
+
20
+
21
+ @app.get("/health")
22
+ async def health():
23
+ return {"status": "healthy"}
ensemble-core/requirements.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ fastapi\nuvicorn\npydantic
ensemble-core/scripts/create_dataset.py ADDED
@@ -0,0 +1,169 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Build a balanced 10k-real / 10k-fake face-image folder with maximum
4
+ deepfake-tech variety.
5
+
6
+ Usage:
7
+ python build_face_dataset.py --out_dir ./faces20k --seed 42
8
+ """
9
+
10
+ import argparse
11
+ import os
12
+ import pathlib
13
+ import random
14
+ import shutil
15
+ import subprocess
16
+ import sys
17
+ import zipfile
18
+ from collections import defaultdict
19
+
20
+ import pandas as pd
21
+ from tqdm import tqdm
22
+
23
+ # ----------------------------------------------------------------------
24
+ # 1. Edit here to add / remove sources
25
+ # ----------------------------------------------------------------------
26
+ DATASETS = [
27
+ {
28
+ "name": "140k",
29
+ "slug": "xhlulu/140k-real-and-fake-faces",
30
+ "subdirs": {"real": "real", "fake": "fake"},
31
+ "fake_label": "stylegan2",
32
+ },
33
+ {
34
+ "name": "deepfake_real",
35
+ "slug": "manjilkarki/deepfake-and-real-images",
36
+ "subdirs": {"real": "real", "fake": "fake"},
37
+ "fake_label": "pggan_stylegan_mix",
38
+ },
39
+ {
40
+ "name": "dfdc_f150",
41
+ "slug": "sciarrilli/dfdc-f150",
42
+ "subdirs": {"real": "real", "fake": "fake"},
43
+ "fake_label": "dfdc_swaps",
44
+ },
45
+ {
46
+ "name": "faceforensics_imgs",
47
+ "slug": "greatgamedota/faceforensics",
48
+ "subdirs": {"real": "real", "fake": "fake"},
49
+ "fake_label": "ffpp_swaps",
50
+ },
51
+ ]
52
+
53
+ TARGET_PER_CLASS = 10_000
54
+ # ----------------------------------------------------------------------
55
+
56
+
57
+ def kaggle_download(slug: str, dest: pathlib.Path) -> pathlib.Path:
58
+ """Download <slug> to dest/. Returns path of the zip."""
59
+ dest.mkdir(parents=True, exist_ok=True)
60
+ zip_path = dest / f"{slug.split('/')[-1]}.zip"
61
+ if zip_path.exists():
62
+ return zip_path
63
+ print(f"Downloading {slug} …")
64
+ subprocess.run(
65
+ ["kaggle", "datasets", "download", "-d", slug, "-p", str(dest), "--quiet"],
66
+ check=True,
67
+ )
68
+ return zip_path
69
+
70
+
71
+ def extract(zip_path: pathlib.Path, dest: pathlib.Path) -> pathlib.Path:
72
+ """Unzip if needed. Returns extraction dir."""
73
+ extract_dir = dest / zip_path.stem
74
+ if extract_dir.exists():
75
+ return extract_dir
76
+ print(f"Extracting {zip_path.name} …")
77
+ with zipfile.ZipFile(zip_path) as zf:
78
+ zf.extractall(path=extract_dir)
79
+ return extract_dir
80
+
81
+
82
+ def glob_images(root: pathlib.Path, pattern: str):
83
+ return list(root.glob(pattern)) + list(root.glob(pattern.replace("jpg", "png")))
84
+
85
+
86
+ def main(out_dir: pathlib.Path, seed: int):
87
+ random.seed(seed)
88
+ temp_root = out_dir / "_raw"
89
+ real_pool, fake_pool = [], []
90
+ fake_source_tag = {} # path -> dataset tag
91
+
92
+ # ------------------------------------------------------------------
93
+ # 2. Pull sources
94
+ # ------------------------------------------------------------------
95
+ for ds in DATASETS:
96
+ zip_path = kaggle_download(ds["slug"], temp_root)
97
+ extract_dir = extract(zip_path, temp_root)
98
+ real_dir = extract_dir / ds["subdirs"]["real"]
99
+ fake_dir = extract_dir / ds["subdirs"]["fake"]
100
+ real_pool += glob_images(real_dir, "**/*.jpg")
101
+ fakes = glob_images(fake_dir, "**/*.jpg")
102
+ fake_pool += fakes
103
+ for fp in fakes:
104
+ fake_source_tag[str(fp)] = ds["fake_label"]
105
+
106
+ # sanity check
107
+ if len(real_pool) < TARGET_PER_CLASS or len(fake_pool) < TARGET_PER_CLASS:
108
+ print("Not enough images – add another dataset.", file=sys.stderr)
109
+ sys.exit(1)
110
+
111
+ # ------------------------------------------------------------------
112
+ # 3. Sample
113
+ # ------------------------------------------------------------------
114
+ random.shuffle(real_pool)
115
+ random.shuffle(fake_pool)
116
+
117
+ # try to spread fake quota equally over sources
118
+ per_source_quota = TARGET_PER_CLASS // len(DATASETS)
119
+ selected_fake = []
120
+ taken = defaultdict(int)
121
+ for fp in fake_pool:
122
+ tag = fake_source_tag[str(fp)]
123
+ if taken[tag] < per_source_quota:
124
+ selected_fake.append(fp)
125
+ taken[tag] += 1
126
+ if len(selected_fake) == TARGET_PER_CLASS:
127
+ break
128
+ # top-up if we’re short (some sets too small)
129
+ if len(selected_fake) < TARGET_PER_CLASS:
130
+ needed = TARGET_PER_CLASS - len(selected_fake)
131
+ selected_fake += fake_pool[len(selected_fake) : len(selected_fake) + needed]
132
+
133
+ selected_real = real_pool[:TARGET_PER_CLASS]
134
+
135
+ # ------------------------------------------------------------------
136
+ # 4. Copy to final tree + manifest
137
+ # ------------------------------------------------------------------
138
+ for cls in ("real", "fake"):
139
+ (out_dir / cls).mkdir(parents=True, exist_ok=True)
140
+
141
+ manifest_rows = []
142
+
143
+ def copy_files(file_list, cls):
144
+ for src in tqdm(file_list, desc=f"Copying {cls}"):
145
+ dst = out_dir / cls / src.name
146
+ shutil.copy(src, dst)
147
+ manifest_rows.append(
148
+ {
149
+ "filename": dst.name,
150
+ "label": cls,
151
+ "source": fake_source_tag.get(str(src), "n/a"),
152
+ }
153
+ )
154
+
155
+ copy_files(selected_real, "real")
156
+ copy_files(selected_fake, "fake")
157
+
158
+ pd.DataFrame(manifest_rows).to_csv(out_dir / "manifest.csv", index=False)
159
+ print("Done β†’", out_dir)
160
+
161
+
162
+ if __name__ == "__main__":
163
+ p = argparse.ArgumentParser()
164
+ p.add_argument(
165
+ "--out_dir", default="faces20k", type=pathlib.Path, help="destination folder"
166
+ )
167
+ p.add_argument("--seed", default=42, type=int)
168
+ args = p.parse_args()
169
+ main(args.out_dir, args.seed)
ensemble-core/scripts/meta_feature_generator.py ADDED
@@ -0,0 +1,366 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ DeepSafe Meta-Feature Generator
4
+ ===============================
5
+
6
+ Orchestrates the generation of meta-feature datasets for training stacking ensembles.
7
+ This component acts as a data ingestion pipeline that:
8
+ 1. Scans a target directory for labeled media (Real/Fake).
9
+ 2. Queries the distributed model microservices to obtain base probability scores.
10
+ 3. Aggregates these scores into a structured feature matrix (CSV) for the meta-learner.
11
+
12
+ Architectural Note:
13
+ This script is designed to be fault-tolerant. If a specific model microservice is unreachable
14
+ or fails for a subset of files, the pipeline continues, recording NaNs for those features.
15
+ This ensures that a single model failure does not halt the entire training data generation process,
16
+ though downstream imputers must handle these missing values.
17
+ """
18
+
19
+ import argparse
20
+ import gc
21
+ import json
22
+ import os
23
+ import sys
24
+ import time
25
+ from typing import Any, Dict, List, Optional
26
+
27
+ import numpy as np
28
+ import pandas as pd
29
+ from rich.console import Console
30
+ from rich.panel import Panel
31
+ from rich.progress import (
32
+ BarColumn,
33
+ MofNCompleteColumn,
34
+ Progress,
35
+ SpinnerColumn,
36
+ TextColumn,
37
+ TimeElapsedColumn,
38
+ )
39
+ from rich.table import Table
40
+
41
+ # Ensure utils is importable regardless of execution context.
42
+ # This fallback is necessary when running the script directly from the project root
43
+ # without an installed package structure.
44
+ try:
45
+ from utils.api_client import APIClient
46
+ from utils.config_manager import ConfigManager
47
+ from utils.media_handler import MediaHandler
48
+ except ImportError:
49
+ project_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
50
+ if project_root not in sys.path:
51
+ sys.path.insert(0, project_root)
52
+ try:
53
+ from utils.api_client import APIClient
54
+ from utils.config_manager import ConfigManager
55
+ from utils.media_handler import MediaHandler
56
+ except ImportError as e:
57
+ print(f"Critical Error: Failed to resolve utils dependency. {e}")
58
+ sys.exit(1)
59
+
60
+
61
+ console = Console(width=120)
62
+
63
+
64
+ class MetaFeatureGenerator:
65
+ """
66
+ Manages the ETL process for meta-learning datasets.
67
+
68
+ Attributes:
69
+ media_type (str): The domain of operation (image, video, audio).
70
+ config_manager (ConfigManager): Centralized configuration handler.
71
+ api_client (APIClient): Interface for communicating with model microservices.
72
+ """
73
+
74
+ def __init__(self, media_type: str, config_manager: ConfigManager):
75
+ self.media_type = media_type
76
+ self.config_manager = config_manager
77
+ # run_from_host=True implies we are running outside the docker network (e.g., local dev),
78
+ # so we use localhost ports mapped in docker-compose.
79
+ self.api_client = APIClient(config_manager, media_type, run_from_host=True)
80
+ self.media_handler = MediaHandler(config_manager)
81
+ self.base_model_names = list(
82
+ config_manager.get_model_endpoints(media_type).keys()
83
+ )
84
+
85
+ if not self.base_model_names:
86
+ console.print(
87
+ f"[bold red]Configuration Error: No base models defined for '{media_type}'.[/bold red]"
88
+ )
89
+ sys.exit(1)
90
+
91
+ def generate(
92
+ self,
93
+ input_dir: str,
94
+ output_csv_path: str,
95
+ default_threshold: float,
96
+ specific_models: Optional[List[str]] = None,
97
+ ):
98
+ """
99
+ Executes the generation pipeline.
100
+
101
+ Args:
102
+ input_dir: Root directory containing 'Real' and 'Fake' subdirectories.
103
+ output_csv_path: Destination for the resulting feature matrix.
104
+ default_threshold: Decision threshold passed to models (mostly for logging/reference).
105
+ specific_models: Optional filter to run only a subset of available models.
106
+ """
107
+
108
+ console.print(
109
+ Panel(
110
+ f"[bold cyan]Meta-Feature Generation Protocol ({self.media_type.capitalize()})[/bold cyan]\n"
111
+ f"Source: {input_dir}\n"
112
+ f"Target: {output_csv_path}\n"
113
+ f"Active Models: {specific_models or 'All configured'}",
114
+ title="Pipeline Configuration",
115
+ border_style="blue",
116
+ expand=False,
117
+ )
118
+ )
119
+
120
+ # Discovery phase: Scan filesystem for valid media files and infer ground truth from directory structure.
121
+ media_files_with_gt = self.media_handler.find_media_files(
122
+ input_dir, self.media_type
123
+ )
124
+ if not media_files_with_gt:
125
+ console.print(
126
+ f"[bold red]Abort: No valid {self.media_type} files found in '{input_dir}'.[/bold red]"
127
+ )
128
+ return
129
+
130
+ # Determine the execution scope (subset of models vs all).
131
+ models_to_query = self.base_model_names
132
+ if specific_models:
133
+ models_to_query = [m for m in specific_models if m in self.base_model_names]
134
+ if not models_to_query:
135
+ console.print(
136
+ f"[bold red]Configuration Mismatch: Requested models {specific_models} are not configured for '{self.media_type}'.[/bold red]"
137
+ )
138
+ return
139
+ console.print(f"Scope restricted to: {models_to_query}")
140
+
141
+ all_feature_data = []
142
+
143
+ # Execution phase: Iterate through files and query models.
144
+ # We use a rich progress bar for observability during long-running batch processes.
145
+ with Progress(
146
+ SpinnerColumn(),
147
+ TextColumn("[progress.description]{task.description}"),
148
+ BarColumn(),
149
+ MofNCompleteColumn(),
150
+ TimeElapsedColumn(),
151
+ ) as progress:
152
+ total_files = len(media_files_with_gt)
153
+ outer_task = progress.add_task(
154
+ f"Processing {self.media_type} corpus...", total=total_files
155
+ )
156
+
157
+ for file_idx, (media_path, ground_truth_label) in enumerate(
158
+ media_files_with_gt
159
+ ):
160
+ media_file_name = os.path.basename(media_path)
161
+ progress.update(
162
+ outer_task,
163
+ description=f"Processing: [cyan]{media_file_name}[/cyan]",
164
+ )
165
+
166
+ # Pre-encode media to base64 once to avoid redundant I/O operations per model.
167
+ encoded_media = self.media_handler.encode_media_to_base64(media_path)
168
+ if not encoded_media:
169
+ console.print(
170
+ f"[yellow]Skip: Encoding failed for {media_file_name}.[/yellow]"
171
+ )
172
+ progress.advance(outer_task)
173
+ continue
174
+
175
+ # Feature vector initialization
176
+ current_media_features: Dict[str, Any] = {
177
+ "media_path": media_path,
178
+ "media_name": media_file_name,
179
+ # Map string labels to numeric binary targets: Fake=1, Real=0.
180
+ "ground_truth": (
181
+ 1
182
+ if ground_truth_label == "Fake"
183
+ else (0 if ground_truth_label == "Real" else -1)
184
+ ),
185
+ }
186
+
187
+ # Initialize feature columns with NaN. This ensures structural consistency in the DataFrame
188
+ # even if specific model queries fail.
189
+ for model_name_cfg in self.base_model_names:
190
+ current_media_features[f"{model_name_cfg}_prob"] = np.nan
191
+
192
+ # Query loop
193
+ for model_name_query in models_to_query:
194
+ model_result = self.api_client.test_with_individual_model(
195
+ model_name_query, media_path, encoded_media, default_threshold
196
+ )
197
+
198
+ if (
199
+ "error" not in model_result
200
+ and model_result.get("probability") is not None
201
+ ):
202
+ current_media_features[f"{model_name_query}_prob"] = (
203
+ model_result["probability"]
204
+ )
205
+ else:
206
+ # Log failure but do not interrupt the pipeline. Robustness is key here.
207
+ error_msg = model_result.get(
208
+ "error", "Invalid response payload"
209
+ )
210
+ console.print(
211
+ f"[yellow]Model Failure: {model_name_query} on {media_file_name}. Reason: {error_msg}.[/yellow]",
212
+ highlight=False,
213
+ )
214
+
215
+ all_feature_data.append(current_media_features)
216
+ progress.advance(outer_task)
217
+
218
+ # Explicit garbage collection to prevent memory bloat during large dataset processing.
219
+ gc.collect()
220
+
221
+ if not all_feature_data:
222
+ console.print(
223
+ "[bold red]Pipeline Failure: No features generated.[/bold red]"
224
+ )
225
+ return
226
+
227
+ # Data serialization and validation
228
+ meta_features_df = pd.DataFrame(all_feature_data)
229
+
230
+ # Filter invalid ground truth (should be handled by discovery, but defensive programming is good).
231
+ meta_features_df = meta_features_df[meta_features_df["ground_truth"] != -1]
232
+
233
+ if meta_features_df.empty:
234
+ console.print(
235
+ "[bold red]Data Error: No valid labeled data remaining after processing.[/bold red]"
236
+ )
237
+ return
238
+
239
+ # Schema enforcement: Ensure all expected columns exist.
240
+ expected_prob_cols = [f"{mn}_prob" for mn in self.base_model_names]
241
+ for col in expected_prob_cols:
242
+ if col not in meta_features_df.columns:
243
+ meta_features_df[col] = np.nan
244
+
245
+ # Column ordering for readability and consistency.
246
+ ordered_prob_cols = sorted(
247
+ [col for col in meta_features_df.columns if col.endswith("_prob")]
248
+ )
249
+ final_cols_order = (
250
+ ["media_path", "media_name"] + ordered_prob_cols + ["ground_truth"]
251
+ )
252
+ meta_features_df = meta_features_df[final_cols_order]
253
+
254
+ try:
255
+ os.makedirs(
256
+ os.path.dirname(os.path.abspath(output_csv_path)), exist_ok=True
257
+ )
258
+ meta_features_df.to_csv(output_csv_path, index=False, float_format="%.6f")
259
+
260
+ console.print(
261
+ f"\n[bold green]Success: Dataset persisted to {os.path.abspath(output_csv_path)}[/bold green]"
262
+ )
263
+ console.print(f"Dimensions: {meta_features_df.shape}")
264
+
265
+ # Quality Assurance: Report missing values to inform downstream handling strategies.
266
+ nan_summary_table = Table(
267
+ title="Data Quality Report (Missing Values)", show_lines=True
268
+ )
269
+ nan_summary_table.add_column("Feature", style="cyan")
270
+ nan_summary_table.add_column(
271
+ "Missing Count", style="magenta", justify="right"
272
+ )
273
+ nan_summary_table.add_column("Missing %", style="yellow", justify="right")
274
+
275
+ for col in ordered_prob_cols:
276
+ nan_count = meta_features_df[col].isnull().sum()
277
+ nan_percent = (
278
+ (nan_count / len(meta_features_df)) * 100
279
+ if len(meta_features_df) > 0
280
+ else 0
281
+ )
282
+ nan_summary_table.add_row(col, str(nan_count), f"{nan_percent:.2f}%")
283
+ console.print(nan_summary_table)
284
+
285
+ except Exception as e:
286
+ console.print(
287
+ f"[bold red]I/O Error: Failed to write output CSV. {e}[/bold red]"
288
+ )
289
+
290
+
291
+ def main():
292
+ parser = argparse.ArgumentParser(
293
+ description="DeepSafe Meta-Feature Generator: ETL for Stacking Ensemble Training Data.",
294
+ formatter_class=argparse.RawDescriptionHelpFormatter,
295
+ )
296
+ parser.add_argument(
297
+ "--media-type",
298
+ type=str,
299
+ choices=["image", "video", "audio"],
300
+ required=True,
301
+ help="Target domain. Defines the model registry subset.",
302
+ )
303
+ parser.add_argument(
304
+ "--input-dir",
305
+ type=str,
306
+ required=True,
307
+ help="Source directory. Must contain 'Real' and 'Fake' subdirectories for label inference.",
308
+ )
309
+ parser.add_argument(
310
+ "--output-csv",
311
+ type=str,
312
+ required=True,
313
+ help="Destination path for the generated feature matrix.",
314
+ )
315
+ parser.add_argument(
316
+ "--threshold",
317
+ type=float,
318
+ help="Decision threshold override (0.0-1.0). Defaults to system config.",
319
+ )
320
+ parser.add_argument(
321
+ "--specific-models",
322
+ type=str,
323
+ help="Optional filter: Comma-separated list of model identifiers to query.",
324
+ )
325
+ parser.add_argument(
326
+ "--config-path", type=str, default=None, help=f"Configuration override path."
327
+ )
328
+
329
+ args = parser.parse_args()
330
+
331
+ # Initialize configuration subsystem
332
+ cfg_manager = ConfigManager(config_path=args.config_path)
333
+ if not cfg_manager.is_config_loaded_successfully():
334
+ sys.exit(1)
335
+
336
+ default_thresh_from_config = cfg_manager.get_default("default_threshold", 0.5)
337
+ query_threshold = (
338
+ args.threshold if args.threshold is not None else default_thresh_from_config
339
+ )
340
+
341
+ specific_models_list = (
342
+ [m.strip() for m in args.specific_models.split(",")]
343
+ if args.specific_models
344
+ else None
345
+ )
346
+
347
+ generator = MetaFeatureGenerator(args.media_type, cfg_manager)
348
+ generator.generate(
349
+ args.input_dir, args.output_csv, query_threshold, specific_models_list
350
+ )
351
+
352
+
353
+ if __name__ == "__main__":
354
+ try:
355
+ main()
356
+ except KeyboardInterrupt:
357
+ console.print("\n[bold yellow]Process Interrupted by User.[/bold yellow]")
358
+ sys.exit(0)
359
+ except Exception as e:
360
+ console.print(f"\n[bold red]Fatal Error: {e}[/bold red]")
361
+ import traceback
362
+
363
+ console.print(
364
+ Panel(traceback.format_exc(), title="Stack Trace", border_style="red")
365
+ )
366
+ sys.exit(1)
ensemble-core/scripts/train_meta_learner_advanced.py ADDED
@@ -0,0 +1,1228 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ DeepSafe Advanced Meta-Learner Training Suite (train_meta_learner_advanced.py)
4
+ ==============================================================================
5
+
6
+ This script trains and evaluates various meta-learners (stacking ensembles)
7
+ for deepfake detection. It takes a CSV file of meta-features (outputs from
8
+ base deepfake detection models) and ground truth labels as input.
9
+
10
+ Key Features:
11
+ -------------
12
+ 1. Modality-Specific Training: Supports training separate meta-learners for
13
+ different media types (image, video, audio) using the `--media-type` argument.
14
+ This ensures that the meta-learner is optimized for the characteristics of
15
+ the base models relevant to that modality.
16
+ 2. Data Preprocessing: Includes imputation for missing values (e.g., if a base
17
+ model failed) and feature scaling.
18
+ 3. Multiple Meta-Learner Models: Trains and evaluates several standard classifiers
19
+ (Logistic Regression, Random Forest, Gradient Boosting, SVC, KNN, Naive Bayes)
20
+ and, if available, advanced models like XGBoost and LightGBM.
21
+ 4. Hyperparameter Optimization:
22
+ - Supports Optuna for efficient hyperparameter search.
23
+ - Falls back to GridSearchCV if Optuna is not installed or if specified.
24
+ 5. Comprehensive Evaluation:
25
+ - Calculates Accuracy, F1-Score, Precision, Recall, and ROC AUC for each model.
26
+ - Generates classification reports and confusion matrices.
27
+ - Plots ROC curves for visual comparison of all trained meta-learners and
28
+ simple ensemble baselines.
29
+ 6. Simple Ensemble Baselines: Also evaluates simple averaging and majority vote
30
+ ensembles for comparison against more complex stacking models. Includes an
31
+ option for optimized weighted averaging.
32
+ 7. Artifact Generation:
33
+ - Saves all trained meta-learner models (e.g., .joblib files).
34
+ - Saves the data preprocessor (imputer + scaler).
35
+ - Saves the list of feature columns used during training.
36
+ - Saves a summary of all experiment metrics in JSON format.
37
+ - The final, best-performing trainable meta-learner and its associated
38
+ preprocessors are saved with generic names inside media-type specific
39
+ subfolders (e.g., api_artifacts_dir/image/deepsafe_meta_learner.joblib).
40
+ 8. Configurable Output: Allows specifying separate directories for general
41
+ experiment outputs and for API-ready deployment artifacts.
42
+
43
+ CLI Usage:
44
+ ----------
45
+ python train_meta_learner_advanced.py \\
46
+ --media-type [image|video|audio] \\
47
+ --meta-file /path/to/meta_features_[media_type].csv \\
48
+ --output-dir ./meta_learning_experiment_runs/ \\
49
+ --api-artifacts-dir ./api/meta_model_artifacts/ \\
50
+ [--optimizer optuna|gridsearch] \\
51
+ [--optuna-trials 50] \\
52
+ [--weights /path/to/custom_weights.json]
53
+
54
+ Arguments:
55
+ ----------
56
+ --media-type {image,video,audio}
57
+ (Required) The type of media for which the meta-learner
58
+ is being trained. This affects output artifact naming.
59
+ --meta-file META_FILE
60
+ (Required) Path to the CSV file containing meta-features
61
+ (base model outputs) and a 'ground_truth' column.
62
+ --output-dir OUTPUT_DIR
63
+ Base directory for saving all experiment-related outputs
64
+ (logs, plots, individual model files from this run).
65
+ A timestamped, media-type-specific subdirectory will be created.
66
+ (Default: ./meta_learning_experiment_runs/)
67
+ --api-artifacts-dir API_ARTIFACTS_DIR
68
+ Directory to save the final, API-ready deployment artifacts
69
+ (e.g., ./api/meta_model_artifacts/image/deepsafe_meta_learner.joblib).
70
+ (Default: ./api/meta_model_artifacts/)
71
+ --optimizer {optuna,gridsearch}
72
+ Hyperparameter optimization strategy (Default: optuna).
73
+ --optuna-trials N
74
+ Number of trials for Optuna optimization (Default: 50).
75
+ --weights WEIGHTS_PATH_OR_JSON
76
+ Optional. Path to a JSON file or a JSON string defining
77
+ custom weights for the 'Provided_Weighted_Average' ensemble.
78
+ Keys should be base model names (without '_prob' suffix).
79
+
80
+ Example (Image Meta-Learner):
81
+ -----------------------------
82
+ python train_meta_learner_advanced.py \\
83
+ --media-type image \\
84
+ --meta-file ./meta_learning_data/meta_features_image.csv \\
85
+ --output-dir ./ml_experiments_images \\
86
+ --api-artifacts-dir ./deepsafe_private/api/meta_model_artifacts \\
87
+ --optimizer optuna \\
88
+ --optuna-trials 100
89
+
90
+ This will train image-specific meta-learners, save experiment details in
91
+ `./ml_experiments_images/experiments_image_YYYYMMDD_HHMMSS/`, and place
92
+ API-ready artifacts like `deepsafe_meta_learner.joblib` into
93
+ `./deepsafe_private/api/meta_model_artifacts/image/`.
94
+ """
95
+
96
+ import argparse
97
+ import itertools
98
+ import json
99
+ import os
100
+ import time
101
+ from typing import Any, Dict, List, Optional, Tuple
102
+
103
+ import joblib
104
+ import matplotlib.pyplot as plt
105
+ import numpy as np
106
+ import pandas as pd
107
+ import seaborn as sns
108
+ from rich.console import Console
109
+ from rich.panel import Panel
110
+ from rich.progress import (
111
+ BarColumn,
112
+ MofNCompleteColumn,
113
+ Progress,
114
+ SpinnerColumn,
115
+ TextColumn,
116
+ TimeElapsedColumn,
117
+ )
118
+ from rich.table import Table
119
+ from sklearn.ensemble import GradientBoostingClassifier, RandomForestClassifier
120
+ from sklearn.impute import SimpleImputer
121
+ from sklearn.linear_model import LogisticRegression
122
+ from sklearn.metrics import (
123
+ accuracy_score,
124
+ auc,
125
+ classification_report,
126
+ confusion_matrix,
127
+ f1_score,
128
+ precision_score,
129
+ recall_score,
130
+ roc_auc_score,
131
+ roc_curve,
132
+ )
133
+ from sklearn.model_selection import StratifiedKFold, train_test_split
134
+ from sklearn.naive_bayes import GaussianNB
135
+ from sklearn.neighbors import KNeighborsClassifier
136
+ from sklearn.pipeline import Pipeline
137
+ from sklearn.preprocessing import StandardScaler
138
+ from sklearn.svm import SVC
139
+
140
+ # --- Optional Advanced Hyperparameter Optimization & Models ---
141
+ OPTIMIZER_CHOICE_DEFAULT = "optuna"
142
+
143
+ try:
144
+ import optuna
145
+
146
+ OPTIMIZER_AVAILABLE_OPTUNA = True
147
+ except ImportError:
148
+ optuna = None
149
+ OPTIMIZER_AVAILABLE_OPTUNA = False
150
+
151
+ from sklearn.model_selection import GridSearchCV
152
+
153
+ try:
154
+ from xgboost import XGBClassifier
155
+
156
+ XGBOOST_AVAILABLE = True
157
+ except ImportError:
158
+ XGBClassifier = None
159
+ XGBOOST_AVAILABLE = False
160
+
161
+ try:
162
+ from lightgbm import LGBMClassifier
163
+
164
+ LIGHTGBM_AVAILABLE = True
165
+ except ImportError:
166
+ LGBMClassifier = None
167
+ LIGHTGBM_AVAILABLE = False
168
+
169
+ console = Console(width=120)
170
+
171
+ # --- Configuration ---
172
+ DEFAULT_EXPERIMENT_OUTPUT_DIR_BASE = "./meta_learning_experiment_runs"
173
+ DEFAULT_API_ARTIFACTS_DIR = "./api/meta_model_artifacts"
174
+ DEFAULT_THRESHOLD_FOR_SIMPLE_ENSEMBLES = 0.5
175
+ N_OPTUNA_TRIALS_DEFAULT = 50
176
+ CV_FOLDS_DEFAULT = 5
177
+
178
+
179
+ # --- Helper Functions ---
180
+ class NpEncoder(json.JSONEncoder):
181
+ def default(self, o: Any) -> Any:
182
+ if isinstance(o, np.integer):
183
+ return int(o)
184
+ if isinstance(o, np.floating):
185
+ return float(o)
186
+ if isinstance(o, np.ndarray):
187
+ return o.tolist()
188
+ return super(NpEncoder, self).default(o)
189
+
190
+
191
+ def evaluate_model_predictions(
192
+ y_true: np.ndarray,
193
+ y_pred_class: np.ndarray,
194
+ y_pred_proba: Optional[np.ndarray],
195
+ model_name: str = "Model",
196
+ ) -> Dict[str, Any]:
197
+ metrics: Dict[str, Any] = {"name": model_name}
198
+ try:
199
+ metrics["accuracy"] = accuracy_score(y_true, y_pred_class)
200
+ metrics["f1_score"] = f1_score(y_true, y_pred_class, zero_division=0)
201
+ metrics["precision"] = precision_score(y_true, y_pred_class, zero_division=0)
202
+ metrics["recall"] = recall_score(y_true, y_pred_class, zero_division=0)
203
+
204
+ roc_auc_val = np.nan
205
+ if y_pred_proba is not None and len(np.unique(y_true)) > 1:
206
+ if not (
207
+ len(np.unique(y_pred_proba)) < 2 and len(y_pred_proba) == len(y_true)
208
+ ):
209
+ try:
210
+ roc_auc_val = roc_auc_score(y_true, y_pred_proba)
211
+ except ValueError:
212
+ pass
213
+ metrics["roc_auc"] = roc_auc_val
214
+
215
+ metrics["classification_report_dict"] = classification_report(
216
+ y_true, y_pred_class, digits=4, zero_division=0, output_dict=True
217
+ )
218
+ metrics["confusion_matrix_list"] = confusion_matrix(
219
+ y_true, y_pred_class
220
+ ).tolist()
221
+ metrics["y_pred_test_classes_list"] = (
222
+ y_pred_class.tolist()
223
+ if isinstance(y_pred_class, np.ndarray)
224
+ else y_pred_class
225
+ )
226
+ metrics["y_prob_test_scores_list"] = (
227
+ y_pred_proba.tolist()
228
+ if y_pred_proba is not None and isinstance(y_pred_proba, np.ndarray)
229
+ else y_pred_proba
230
+ )
231
+ except Exception as e:
232
+ console.print(
233
+ f"[bold red]Error during evaluation for {model_name}: {e}[/bold red]"
234
+ )
235
+ for m_key in ["accuracy", "f1_score", "precision", "recall", "roc_auc"]:
236
+ metrics[m_key] = np.nan
237
+ metrics["classification_report_dict"] = {}
238
+ metrics["confusion_matrix_list"] = []
239
+ return metrics
240
+
241
+
242
+ def plot_roc_curves_all(
243
+ experiment_results_dict: Dict[str, Dict[str, Any]],
244
+ y_true_labels: np.ndarray,
245
+ output_dir_path: str,
246
+ media_type: str,
247
+ ):
248
+ plt.figure(figsize=(12, 10))
249
+ plot_count = 0
250
+ for model_key, result_data in experiment_results_dict.items():
251
+ if (
252
+ "y_prob_test_scores_list" in result_data
253
+ and result_data["y_prob_test_scores_list"] is not None
254
+ ):
255
+ proba_scores = np.array(result_data["y_prob_test_scores_list"])
256
+ if len(np.unique(y_true_labels)) < 2 or (
257
+ proba_scores.ndim > 0
258
+ and len(np.unique(proba_scores)) < 2
259
+ and len(proba_scores) == len(y_true_labels)
260
+ ):
261
+ continue
262
+ try:
263
+ fpr, tpr, _ = roc_curve(y_true_labels, proba_scores)
264
+ roc_auc_value = result_data.get("roc_auc", auc(fpr, tpr))
265
+ if pd.notna(roc_auc_value):
266
+ plt.plot(
267
+ fpr,
268
+ tpr,
269
+ lw=1.8,
270
+ label=f"{model_key} (AUC = {roc_auc_value:.4f})",
271
+ )
272
+ plot_count += 1
273
+ except ValueError as e:
274
+ console.print(
275
+ f"[yellow]Could not plot ROC for {model_key} ({media_type}): {e}[/yellow]"
276
+ )
277
+
278
+ if plot_count > 0:
279
+ plt.plot([0, 1], [0, 1], color="grey", lw=1.5, linestyle="--")
280
+ plt.xlim([-0.01, 1.0])
281
+ plt.ylim([0.0, 1.01])
282
+ plt.xlabel("False Positive Rate", fontsize=13)
283
+ plt.ylabel("True Positive Rate", fontsize=13)
284
+ plt.title(
285
+ f"Meta-Learner & Ensemble ROC Curves ({media_type.capitalize()})",
286
+ fontsize=15,
287
+ )
288
+ plt.legend(loc="lower right", fontsize="small", frameon=True)
289
+ plt.grid(alpha=0.35, linestyle=":")
290
+ plt.tight_layout()
291
+ plot_path = os.path.join(
292
+ output_dir_path, f"all_meta_learners_roc_curves_{media_type}.png"
293
+ )
294
+ plt.savefig(plot_path, dpi=150)
295
+ console.print(
296
+ f"Combined ROC curves plot for {media_type} saved to [green]{plot_path}[/green]"
297
+ )
298
+ else:
299
+ console.print(f"[yellow]No valid ROC curves to plot for {media_type}.[/yellow]")
300
+ plt.close()
301
+
302
+
303
+ def optimize_average_weights_simple_grid(
304
+ X_val_probs: np.ndarray,
305
+ y_val_true: np.ndarray,
306
+ num_base_models: int,
307
+ weight_options: Optional[List[float]] = None,
308
+ ) -> np.ndarray:
309
+ if weight_options is None:
310
+ weight_options = [0.25, 0.5, 0.75, 1.0, 1.25, 1.5, 1.75, 2.0]
311
+ best_auc_val = -1.0
312
+ best_weights_val = np.ones(num_base_models)
313
+
314
+ max_combinations_exhaustive = 5**4
315
+ num_random_samples_if_large = 2000
316
+
317
+ if num_base_models <= 0:
318
+ console.print(
319
+ "[yellow]No base models to optimize weights for. Returning default weights.[/yellow]"
320
+ )
321
+ return best_weights_val
322
+
323
+ if (
324
+ num_base_models <= 4
325
+ and len(weight_options) ** num_base_models <= max_combinations_exhaustive
326
+ ):
327
+ weight_candidates = list(
328
+ itertools.product(weight_options, repeat=num_base_models)
329
+ )
330
+ console.print(
331
+ f"Optimizing average weights with exhaustive grid search ({len(weight_candidates)} trials)."
332
+ )
333
+ else:
334
+ console.print(
335
+ f"[yellow]Optimizing average weights with random sampling ({num_random_samples_if_large} trials due to {num_base_models} models).[/yellow]"
336
+ )
337
+ weight_candidates = [
338
+ np.array(np.random.choice(weight_options, num_base_models))
339
+ for _ in range(num_random_samples_if_large)
340
+ ]
341
+
342
+ with Progress(
343
+ SpinnerColumn(),
344
+ TextColumn("[progress.description]{task.description}"),
345
+ BarColumn(),
346
+ TextColumn("{task.percentage:>3.1f}%"),
347
+ TimeElapsedColumn(),
348
+ MofNCompleteColumn(),
349
+ ) as progress:
350
+ task = progress.add_task("Weight Grid Search", total=len(weight_candidates))
351
+ for current_weights_tuple in weight_candidates:
352
+ current_weights = np.array(current_weights_tuple)
353
+ if np.sum(current_weights) == 0:
354
+ progress.update(task, advance=1)
355
+ continue
356
+
357
+ if X_val_probs.shape[0] == 0:
358
+ progress.update(task, advance=1)
359
+ continue
360
+ weighted_avg_probs_val_set = np.average(
361
+ X_val_probs, axis=1, weights=current_weights
362
+ )
363
+
364
+ current_auc_val = 0.0
365
+ if len(np.unique(y_val_true)) > 1 and not (
366
+ len(np.unique(weighted_avg_probs_val_set)) < 2
367
+ and len(weighted_avg_probs_val_set) == len(y_val_true)
368
+ ):
369
+ try:
370
+ current_auc_val = roc_auc_score(
371
+ y_val_true, weighted_avg_probs_val_set
372
+ )
373
+ except ValueError:
374
+ pass
375
+ if current_auc_val > best_auc_val:
376
+ best_auc_val, best_weights_val = current_auc_val, current_weights
377
+ progress.update(task, advance=1)
378
+
379
+ console.print(
380
+ f"Best weights from validation grid search: {best_weights_val.tolist()} with Val AUC: {best_auc_val:.4f}"
381
+ )
382
+ return best_weights_val
383
+
384
+
385
+ # --- Main Experimentation Function ---
386
+ def run_meta_learning_experiments(
387
+ meta_features_file: str,
388
+ output_dir_base: str,
389
+ api_artifacts_dir: str,
390
+ media_type: str,
391
+ optimizer_type: str,
392
+ n_optuna_trials_config: int,
393
+ provided_custom_weights: Optional[Dict[str, float]] = None,
394
+ ):
395
+ global OPTIMIZER_CHOICE, N_OPTUNA_TRIALS
396
+ OPTIMIZER_CHOICE = optimizer_type
397
+ N_OPTUNA_TRIALS = n_optuna_trials_config
398
+
399
+ if OPTIMIZER_CHOICE == "optuna" and not OPTIMIZER_AVAILABLE_OPTUNA:
400
+ console.print(
401
+ "[yellow]Optuna chosen but not installed. Falling back to GridSearchCV.[/yellow]"
402
+ )
403
+ OPTIMIZER_CHOICE = "gridsearch"
404
+
405
+ experiment_run_output_dir = os.path.join(
406
+ output_dir_base, f"experiments_{media_type}_{time.strftime('%Y%m%d_%H%M%S')}"
407
+ )
408
+ os.makedirs(experiment_run_output_dir, exist_ok=True)
409
+
410
+ # Main API artifacts directory (parent for media-specific subfolders)
411
+ os.makedirs(api_artifacts_dir, exist_ok=True)
412
+ # Media-type specific subdirectory within the main api_artifacts_dir
413
+ media_type_api_artifacts_subdir = os.path.join(api_artifacts_dir, media_type)
414
+ os.makedirs(media_type_api_artifacts_subdir, exist_ok=True)
415
+
416
+ console.rule(
417
+ f"[bold cyan]DeepSafe Meta-Learning: {media_type.upper()} (Optimizer: {OPTIMIZER_CHOICE})[/bold cyan]"
418
+ )
419
+ console.print(
420
+ Panel(
421
+ f"Meta-features: {meta_features_file}\n"
422
+ f"Experiment outputs: {os.path.abspath(experiment_run_output_dir)}\n"
423
+ f"API artifacts subfolder: {os.path.abspath(media_type_api_artifacts_subdir)}",
424
+ title="Paths",
425
+ border_style="dim blue",
426
+ expand=False,
427
+ )
428
+ )
429
+ all_experiment_results: Dict[str, Dict[str, Any]] = {}
430
+
431
+ console.rule("[bold]1. Data Loading and Preprocessing[/bold]")
432
+ try:
433
+ df_meta = pd.read_csv(meta_features_file)
434
+ console.print(
435
+ f"Loaded {media_type} meta-features from: [cyan]{meta_features_file}[/cyan], shape: {df_meta.shape}"
436
+ )
437
+ except Exception as e:
438
+ console.print(
439
+ f"[bold red]Fatal Error: Could not load meta-features file: {e}[/bold red]"
440
+ )
441
+ return
442
+
443
+ base_model_prob_features = sorted(
444
+ [col for col in df_meta.columns if col.endswith("_prob")]
445
+ )
446
+ if not base_model_prob_features:
447
+ console.print(
448
+ "[bold red]Fatal Error: No base model probability columns (ending with '_prob') found in CSV.[/bold red]"
449
+ )
450
+ return
451
+
452
+ console.print(
453
+ f"Identified [magenta]{len(base_model_prob_features)}[/magenta] base model probability features: {base_model_prob_features}"
454
+ )
455
+
456
+ temp_exp_feature_cols_path = os.path.join(
457
+ experiment_run_output_dir, f"experiment_feature_columns_{media_type}.json"
458
+ )
459
+ with open(temp_exp_feature_cols_path, "w") as f:
460
+ json.dump(base_model_prob_features, f, indent=2)
461
+
462
+ X_meta_all = df_meta[base_model_prob_features].copy()
463
+ y_meta_all = df_meta["ground_truth"]
464
+
465
+ cols_to_drop_all_nan = X_meta_all.columns[X_meta_all.isnull().all()].tolist()
466
+ if cols_to_drop_all_nan:
467
+ console.print(
468
+ f"[yellow]Warning: Dropping fully NaN columns: {cols_to_drop_all_nan}[/yellow]"
469
+ )
470
+ X_meta_all = X_meta_all.drop(columns=cols_to_drop_all_nan)
471
+ base_model_prob_features = [
472
+ col for col in base_model_prob_features if col not in cols_to_drop_all_nan
473
+ ]
474
+ if not base_model_prob_features:
475
+ console.print(
476
+ "[bold red]Fatal Error: All features became NaN after dropping some columns.[/bold red]"
477
+ )
478
+ return
479
+ with open(temp_exp_feature_cols_path, "w") as f:
480
+ json.dump(base_model_prob_features, f, indent=2)
481
+
482
+ X_meta_train_val, X_meta_test, y_meta_train_val, y_meta_test = train_test_split(
483
+ X_meta_all,
484
+ y_meta_all,
485
+ test_size=0.25,
486
+ random_state=42,
487
+ stratify=y_meta_all if len(np.unique(y_meta_all)) > 1 else None,
488
+ )
489
+ console.print(
490
+ f"Data split: Meta-Train/Val shape {X_meta_train_val.shape}, Meta-Test shape {X_meta_test.shape}"
491
+ )
492
+
493
+ ml_preprocessor = Pipeline(
494
+ [("imputer", SimpleImputer(strategy="median")), ("scaler", StandardScaler())]
495
+ )
496
+ X_meta_train_val_processed = ml_preprocessor.fit_transform(X_meta_train_val)
497
+ X_meta_test_processed = ml_preprocessor.transform(X_meta_test)
498
+
499
+ joblib.dump(
500
+ ml_preprocessor,
501
+ os.path.join(
502
+ experiment_run_output_dir, f"experiment_ml_preprocessor_{media_type}.joblib"
503
+ ),
504
+ )
505
+ console.print(
506
+ f"ML preprocessor for {media_type} (imputer + scaler) fitted and saved for this run."
507
+ )
508
+
509
+ imputer_for_simple_ensembles = ml_preprocessor.named_steps["imputer"]
510
+ X_meta_test_imputed_only_df = pd.DataFrame(
511
+ imputer_for_simple_ensembles.transform(X_meta_test), columns=X_meta_test.columns
512
+ )
513
+
514
+ console.rule("[bold]2. Defining ML Meta-Learners and Hyperparameter Spaces[/bold]")
515
+ models_and_param_spaces: Dict[str, Tuple[Any, Dict[str, Any]]] = {
516
+ "LogisticRegression": (
517
+ LogisticRegression(
518
+ solver="liblinear",
519
+ random_state=42,
520
+ class_weight="balanced",
521
+ max_iter=3000,
522
+ ),
523
+ {
524
+ "C": (
525
+ (0.01, 1000.0, "loguniform")
526
+ if OPTIMIZER_CHOICE == "optuna"
527
+ else [0.01, 0.1, 1, 10, 100, 500]
528
+ )
529
+ },
530
+ ),
531
+ "RandomForest": (
532
+ RandomForestClassifier(random_state=42, class_weight="balanced"),
533
+ {
534
+ "n_estimators": (
535
+ (100, 500, "int")
536
+ if OPTIMIZER_CHOICE == "optuna"
537
+ else [100, 200, 300, 400]
538
+ ),
539
+ "max_depth": (
540
+ (5, 25, "int", True)
541
+ if OPTIMIZER_CHOICE == "optuna"
542
+ else [5, 10, 15, 20, None]
543
+ ),
544
+ "min_samples_split": (
545
+ (2, 20, "int") if OPTIMIZER_CHOICE == "optuna" else [2, 5, 10, 15]
546
+ ),
547
+ "min_samples_leaf": (
548
+ (1, 15, "int") if OPTIMIZER_CHOICE == "optuna" else [1, 5, 10, 15]
549
+ ),
550
+ },
551
+ ),
552
+ "GradientBoosting": (
553
+ GradientBoostingClassifier(random_state=42),
554
+ {
555
+ "n_estimators": (
556
+ (100, 500, "int")
557
+ if OPTIMIZER_CHOICE == "optuna"
558
+ else [100, 200, 300, 400]
559
+ ),
560
+ "learning_rate": (
561
+ (0.005, 0.2, "loguniform")
562
+ if OPTIMIZER_CHOICE == "optuna"
563
+ else [0.01, 0.05, 0.1, 0.15]
564
+ ),
565
+ "max_depth": (
566
+ (3, 10, "int") if OPTIMIZER_CHOICE == "optuna" else [3, 5, 7, 9]
567
+ ),
568
+ },
569
+ ),
570
+ "SVC_Linear": (
571
+ SVC(
572
+ kernel="linear",
573
+ probability=True,
574
+ random_state=42,
575
+ class_weight="balanced",
576
+ max_iter=10000,
577
+ ),
578
+ {
579
+ "C": (
580
+ (0.01, 100.0, "loguniform")
581
+ if OPTIMIZER_CHOICE == "optuna"
582
+ else [0.1, 1, 10, 100]
583
+ )
584
+ },
585
+ ),
586
+ "KNeighbors": (
587
+ KNeighborsClassifier(),
588
+ {
589
+ "n_neighbors": (
590
+ (3, 25, "int", False, 2)
591
+ if OPTIMIZER_CHOICE == "optuna"
592
+ else [3, 5, 7, 11, 15, 19, 23]
593
+ ),
594
+ "weights": (
595
+ (["uniform", "distance"], "categorical")
596
+ if OPTIMIZER_CHOICE == "optuna"
597
+ else ["uniform", "distance"]
598
+ ),
599
+ },
600
+ ),
601
+ "GaussianNB": (GaussianNB(), {}),
602
+ }
603
+ if XGBOOST_AVAILABLE and XGBClassifier:
604
+ models_and_param_spaces["XGBoost"] = (
605
+ XGBClassifier(random_state=42, eval_metric="auc"),
606
+ {
607
+ "n_estimators": (
608
+ (100, 600, "int")
609
+ if OPTIMIZER_CHOICE == "optuna"
610
+ else [100, 200, 300, 400, 500]
611
+ ),
612
+ "learning_rate": (
613
+ (0.005, 0.2, "loguniform")
614
+ if OPTIMIZER_CHOICE == "optuna"
615
+ else [0.01, 0.05, 0.1]
616
+ ),
617
+ "max_depth": (
618
+ (3, 12, "int") if OPTIMIZER_CHOICE == "optuna" else [3, 5, 7, 9, 11]
619
+ ),
620
+ "scale_pos_weight": (
621
+ (
622
+ (np.sum(y_meta_train_val == 0) / np.sum(y_meta_train_val == 1))
623
+ if np.sum(y_meta_train_val == 1) > 0
624
+ else 1.0
625
+ ),
626
+ ),
627
+ },
628
+ )
629
+ if LIGHTGBM_AVAILABLE and LGBMClassifier:
630
+ models_and_param_spaces["LightGBM"] = (
631
+ LGBMClassifier(
632
+ random_state=42, class_weight="balanced", metric="auc", verbosity=-1
633
+ ),
634
+ {
635
+ "n_estimators": (
636
+ (100, 600, "int")
637
+ if OPTIMIZER_CHOICE == "optuna"
638
+ else [100, 200, 300, 400, 500]
639
+ ),
640
+ "learning_rate": (
641
+ (0.005, 0.2, "loguniform")
642
+ if OPTIMIZER_CHOICE == "optuna"
643
+ else [0.01, 0.05, 0.1]
644
+ ),
645
+ "num_leaves": (
646
+ (20, 150, "int")
647
+ if OPTIMIZER_CHOICE == "optuna"
648
+ else [31, 50, 70, 100, 130]
649
+ ),
650
+ },
651
+ )
652
+
653
+ console.rule(
654
+ f"[bold]3. Training and Evaluating ML-based Meta-Learners ({media_type.capitalize()} Stacking)[/bold]"
655
+ )
656
+ cv_strategy = StratifiedKFold(
657
+ n_splits=CV_FOLDS_DEFAULT, shuffle=True, random_state=42
658
+ )
659
+ trained_ml_model_objects: Dict[str, Any] = {}
660
+
661
+ for model_name_key, (
662
+ model_instance_template,
663
+ param_def,
664
+ ) in models_and_param_spaces.items():
665
+ console.rule(
666
+ f"[bold blue]Optimizing & Training {media_type.capitalize()} Meta-Learner: {model_name_key}[/bold blue]",
667
+ style="blue",
668
+ )
669
+ start_train_time = time.time()
670
+ best_estimator_for_model = None
671
+
672
+ if not param_def:
673
+ model_instance_template.fit(X_meta_train_val_processed, y_meta_train_val)
674
+ best_estimator_for_model = model_instance_template
675
+ console.print(
676
+ f"{model_name_key} fitted directly (no hyperparameters tuned)."
677
+ )
678
+ elif OPTIMIZER_CHOICE == "optuna" and optuna:
679
+
680
+ def optuna_objective(trial: optuna.Trial):
681
+ current_params = {}
682
+ for p_name, p_opts in param_def.items():
683
+ if isinstance(p_opts, tuple) and len(p_opts) >= 2:
684
+ suggestion_type_or_values = (
685
+ p_opts[1]
686
+ if p_name == "weights" and p_opts[1] == "categorical"
687
+ else p_opts[2]
688
+ )
689
+ if suggestion_type_or_values == "loguniform":
690
+ current_params[p_name] = trial.suggest_float(
691
+ p_name, p_opts[0], p_opts[1], log=True
692
+ )
693
+ elif suggestion_type_or_values == "uniform":
694
+ current_params[p_name] = trial.suggest_float(
695
+ p_name, p_opts[0], p_opts[1]
696
+ )
697
+ elif suggestion_type_or_values == "int":
698
+ low, high = p_opts[0], p_opts[1]
699
+ can_be_none = p_opts[3] if len(p_opts) > 3 else False
700
+ step = p_opts[4] if len(p_opts) > 4 else 1
701
+ val = trial.suggest_int(p_name, low, high, step=step)
702
+ if can_be_none and trial.suggest_categorical(
703
+ f"{p_name}_use_none", [True, False]
704
+ ):
705
+ val = None
706
+ current_params[p_name] = val
707
+ elif suggestion_type_or_values == "categorical":
708
+ current_params[p_name] = trial.suggest_categorical(
709
+ p_name, p_opts[0]
710
+ )
711
+ elif len(p_opts) == 1 and not isinstance(p_opts[0], list):
712
+ current_params[p_name] = p_opts[0]
713
+ else:
714
+ console.print(
715
+ f"[red]Warning: Unknown Optuna parameter definition for {p_name}: {p_opts}[/red]"
716
+ )
717
+ else:
718
+ if (
719
+ p_name in model_instance_template.get_params()
720
+ and not isinstance(p_opts, tuple)
721
+ ):
722
+ current_params[p_name] = p_opts
723
+
724
+ model_trial = model_instance_template.__class__(
725
+ **model_instance_template.get_params()
726
+ )
727
+ valid_model_params = model_trial.get_params().keys()
728
+ filtered_current_params = {
729
+ k: v for k, v in current_params.items() if k in valid_model_params
730
+ }
731
+ model_trial.set_params(**filtered_current_params)
732
+
733
+ scores = []
734
+ for train_idx, val_idx in cv_strategy.split(
735
+ X_meta_train_val_processed, y_meta_train_val
736
+ ):
737
+ X_fold_train, X_fold_val = (
738
+ X_meta_train_val_processed[train_idx],
739
+ X_meta_train_val_processed[val_idx],
740
+ )
741
+ y_fold_train, y_fold_val = (
742
+ y_meta_train_val.iloc[train_idx],
743
+ y_meta_train_val.iloc[val_idx],
744
+ )
745
+ model_trial.fit(X_fold_train, y_fold_train)
746
+ if hasattr(model_trial, "predict_proba"):
747
+ try:
748
+ y_val_pred_proba = model_trial.predict_proba(X_fold_val)[
749
+ :, 1
750
+ ]
751
+ if len(np.unique(y_fold_val)) < 2 or (
752
+ len(np.unique(y_val_pred_proba)) < 2
753
+ and len(y_val_pred_proba) == len(y_fold_val)
754
+ ):
755
+ scores.append(0.5)
756
+ else:
757
+ scores.append(
758
+ roc_auc_score(y_fold_val, y_val_pred_proba)
759
+ )
760
+ except Exception:
761
+ scores.append(0.0)
762
+ else:
763
+ scores.append(
764
+ f1_score(
765
+ y_fold_val,
766
+ model_trial.predict(X_fold_val),
767
+ zero_division=0,
768
+ )
769
+ )
770
+ return np.mean(scores)
771
+
772
+ study = optuna.create_study(
773
+ direction="maximize", pruner=optuna.pruners.MedianPruner()
774
+ )
775
+ study.optimize(
776
+ optuna_objective,
777
+ n_trials=N_OPTUNA_TRIALS,
778
+ show_progress_bar=True,
779
+ gc_after_trial=True,
780
+ )
781
+
782
+ sklearn_best_params = {}
783
+ for p_name_orig_def, p_opts_def in param_def.items():
784
+ if p_name_orig_def in study.best_params:
785
+ sklearn_best_params[p_name_orig_def] = study.best_params[
786
+ p_name_orig_def
787
+ ]
788
+ if len(p_opts_def) > 3 and p_opts_def[3] is True:
789
+ if (
790
+ study.best_params.get(f"{p_name_orig_def}_use_none", False)
791
+ is True
792
+ ):
793
+ sklearn_best_params[p_name_orig_def] = None
794
+ console.print(
795
+ f"Best Optuna params for {model_name_key} ({media_type}): {sklearn_best_params}"
796
+ )
797
+ best_estimator_for_model = model_instance_template.__class__(
798
+ **model_instance_template.get_params()
799
+ )
800
+ best_estimator_for_model.set_params(**sklearn_best_params)
801
+ best_estimator_for_model.fit(X_meta_train_val_processed, y_meta_train_val)
802
+ else:
803
+ grid_search = GridSearchCV(
804
+ model_instance_template,
805
+ param_def,
806
+ cv=cv_strategy,
807
+ scoring="roc_auc",
808
+ n_jobs=-1,
809
+ verbose=0,
810
+ )
811
+ grid_search.fit(X_meta_train_val_processed, y_meta_train_val)
812
+ best_estimator_for_model = grid_search.best_estimator_
813
+ console.print(
814
+ f"Best GridSearchCV params for {model_name_key} ({media_type}): {grid_search.best_params_}"
815
+ )
816
+
817
+ joblib.dump(
818
+ best_estimator_for_model,
819
+ os.path.join(
820
+ experiment_run_output_dir,
821
+ f"{model_name_key}_meta_learner_{media_type}.joblib",
822
+ ),
823
+ )
824
+ trained_ml_model_objects[model_name_key] = best_estimator_for_model
825
+
826
+ y_test_pred_classes = best_estimator_for_model.predict(X_meta_test_processed)
827
+ y_test_pred_probas = (
828
+ best_estimator_for_model.predict_proba(X_meta_test_processed)[:, 1]
829
+ if hasattr(best_estimator_for_model, "predict_proba")
830
+ else None
831
+ )
832
+ metrics_results = evaluate_model_predictions(
833
+ y_meta_test.values, y_test_pred_classes, y_test_pred_probas, model_name_key
834
+ )
835
+ all_experiment_results[model_name_key] = metrics_results
836
+ train_time = time.time() - start_train_time
837
+ console.print(
838
+ f"[bold]{model_name_key} Test Set Perf. ({media_type}):[/bold] AUC: {metrics_results.get('roc_auc', np.nan):.4f}, F1: {metrics_results.get('f1_score', np.nan):.4f}, Acc: {metrics_results.get('accuracy', np.nan):.4f} (Train time: {train_time:.2f}s)"
839
+ )
840
+
841
+ console.rule(
842
+ f"[bold]4. Evaluating Simple Ensemble Baselines ({media_type.capitalize()} Meta-Test Set)[/bold]"
843
+ )
844
+ avg_probs_meta_test = X_meta_test_imputed_only_df.mean(axis=1).values
845
+ avg_preds_meta_test_classes = (
846
+ avg_probs_meta_test >= DEFAULT_THRESHOLD_FOR_SIMPLE_ENSEMBLES
847
+ ).astype(int)
848
+ all_experiment_results["Simple_Average_Prob"] = evaluate_model_predictions(
849
+ y_meta_test.values,
850
+ avg_preds_meta_test_classes,
851
+ avg_probs_meta_test,
852
+ "Simple_Average_Prob",
853
+ )
854
+ console.print(
855
+ f"[bold]Simple Average Prob Test ({media_type}):[/bold] AUC: {all_experiment_results['Simple_Average_Prob'].get('roc_auc', np.nan):.4f}, F1: {all_experiment_results['Simple_Average_Prob'].get('f1_score', np.nan):.4f}"
856
+ )
857
+
858
+ binarized_X_meta_test = (
859
+ X_meta_test_imputed_only_df.values >= DEFAULT_THRESHOLD_FOR_SIMPLE_ENSEMBLES
860
+ ).astype(int)
861
+ num_models_for_vote = X_meta_test_imputed_only_df.shape[1]
862
+ fake_votes_per_item_meta_test = binarized_X_meta_test.sum(axis=1)
863
+ maj_vote_preds_meta_test_classes = (
864
+ fake_votes_per_item_meta_test >= (num_models_for_vote / 2.0)
865
+ ).astype(int)
866
+ maj_vote_prob_scores_meta_test = (
867
+ fake_votes_per_item_meta_test / num_models_for_vote
868
+ if num_models_for_vote > 0
869
+ else np.full_like(fake_votes_per_item_meta_test, 0.5, dtype=float)
870
+ )
871
+ all_experiment_results["Simple_Majority_Vote"] = evaluate_model_predictions(
872
+ y_meta_test.values,
873
+ maj_vote_preds_meta_test_classes,
874
+ maj_vote_prob_scores_meta_test,
875
+ "Simple_Majority_Vote",
876
+ )
877
+ console.print(
878
+ f"[bold]Simple Majority Vote Test ({media_type}):[/bold] AUC: {all_experiment_results['Simple_Majority_Vote'].get('roc_auc', np.nan):.4f}, F1: {all_experiment_results['Simple_Majority_Vote'].get('f1_score', np.nan):.4f}"
879
+ )
880
+
881
+ if provided_custom_weights:
882
+ current_weights_values = [
883
+ provided_custom_weights.get(fc.replace("_prob", ""), 1.0)
884
+ for fc in base_model_prob_features
885
+ ]
886
+ current_weights_array = np.array(current_weights_values)
887
+
888
+ if (
889
+ len(current_weights_array) == X_meta_test_imputed_only_df.shape[1]
890
+ and np.sum(current_weights_array) > 0
891
+ ):
892
+ prov_weighted_avg_probs_meta_test = np.average(
893
+ X_meta_test_imputed_only_df.values,
894
+ axis=1,
895
+ weights=current_weights_array,
896
+ )
897
+ prov_weighted_avg_preds_meta_test_classes = (
898
+ prov_weighted_avg_probs_meta_test
899
+ >= DEFAULT_THRESHOLD_FOR_SIMPLE_ENSEMBLES
900
+ ).astype(int)
901
+ all_experiment_results["Provided_Weighted_Average"] = (
902
+ evaluate_model_predictions(
903
+ y_meta_test.values,
904
+ prov_weighted_avg_preds_meta_test_classes,
905
+ prov_weighted_avg_probs_meta_test,
906
+ "Provided_Weighted_Average",
907
+ )
908
+ )
909
+ console.print(
910
+ f"[bold]Provided Weighted Average Test ({media_type}):[/bold] AUC: {all_experiment_results['Provided_Weighted_Average'].get('roc_auc', np.nan):.4f}, F1: {all_experiment_results['Provided_Weighted_Average'].get('f1_score', np.nan):.4f}"
911
+ )
912
+ else:
913
+ console.print(
914
+ f"[yellow]Warning: Mismatch in provided_custom_weights keys vs. features for {media_type}, or sum of weights is zero. Skipping.[/yellow]"
915
+ )
916
+
917
+ X_train_val_imputed_for_opt_df = pd.DataFrame(
918
+ ml_preprocessor.named_steps["imputer"].transform(X_meta_train_val),
919
+ columns=base_model_prob_features,
920
+ )
921
+ stratify_opt_split = (
922
+ y_meta_train_val if len(np.unique(y_meta_train_val)) > 1 else None
923
+ )
924
+ X_opt_train_df, X_opt_val_df, y_opt_train_series, y_opt_val_series = (
925
+ train_test_split(
926
+ X_train_val_imputed_for_opt_df,
927
+ y_meta_train_val,
928
+ test_size=0.33,
929
+ random_state=123,
930
+ stratify=stratify_opt_split,
931
+ )
932
+ )
933
+ if X_opt_val_df.shape[0] > 10 and X_opt_val_df.shape[1] > 0:
934
+ console.print(
935
+ f"Optimizing weights for averaging ({media_type}) using a validation split of meta-train data..."
936
+ )
937
+ optimized_avg_weights = optimize_average_weights_simple_grid(
938
+ X_opt_val_df.values, y_opt_val_series.values, X_opt_val_df.shape[1]
939
+ )
940
+ opt_w_avg_probs_meta_test = np.average(
941
+ X_meta_test_imputed_only_df.values, axis=1, weights=optimized_avg_weights
942
+ )
943
+ opt_w_avg_preds_meta_test_classes = (
944
+ opt_w_avg_probs_meta_test >= DEFAULT_THRESHOLD_FOR_SIMPLE_ENSEMBLES
945
+ ).astype(int)
946
+ all_experiment_results["Optimized_Grid_Weighted_Average"] = (
947
+ evaluate_model_predictions(
948
+ y_meta_test.values,
949
+ opt_w_avg_preds_meta_test_classes,
950
+ opt_w_avg_probs_meta_test,
951
+ "Optimized_Grid_Weighted_Average",
952
+ )
953
+ )
954
+ console.print(
955
+ f"[bold]Optimized Grid Weighted Average Test ({media_type}):[/bold] AUC: {all_experiment_results['Optimized_Grid_Weighted_Average'].get('roc_auc', np.nan):.4f}, F1: {all_experiment_results['Optimized_Grid_Weighted_Average'].get('f1_score', np.nan):.4f}"
956
+ )
957
+
958
+ # Save optimized weights to media-type specific subdirectory with generic name
959
+ # (or keep media_type in name if preferred, but API loads generic name from subdir)
960
+ # opt_weights_api_path_generic = os.path.join(media_type_api_artifacts_subdir, "optimized_grid_average_weights.json")
961
+ # For now, keeping the original behavior of saving to main api_artifacts_dir with media_type in name
962
+ opt_weights_api_path_typed = os.path.join(
963
+ api_artifacts_dir, f"optimized_grid_average_weights_{media_type}.json"
964
+ )
965
+ with open(opt_weights_api_path_typed, "w") as f:
966
+ json.dump(
967
+ {
968
+ feat: w
969
+ for feat, w in zip(base_model_prob_features, optimized_avg_weights)
970
+ },
971
+ f,
972
+ indent=2,
973
+ )
974
+ console.print(
975
+ f"Optimized weights for {media_type} saved to API artifacts: [green]{opt_weights_api_path_typed}[/green]"
976
+ )
977
+ else:
978
+ console.print(
979
+ f"[yellow]Validation set for weight optimization ({media_type}) too small or no features. Skipping.[/yellow]"
980
+ )
981
+
982
+ console.rule(
983
+ f"[bold green]5. Overall Experiment Summary & Artifacts ({media_type.capitalize()})[/bold green]"
984
+ )
985
+ summary_table = Table(
986
+ title=f"Meta-Learner & Simple Ensemble Experiment Summary ({media_type.capitalize()} Meta-Test Set)"
987
+ )
988
+ summary_table.add_column(
989
+ "Method/Model", style="cyan", overflow="fold", max_width=35
990
+ )
991
+ summary_table.add_column("Test AUC", style="magenta")
992
+ summary_table.add_column("Test F1", style="green")
993
+ summary_table.add_column("Test Acc.", style="blue")
994
+ summary_table.add_column("Test Prec.", style="yellow")
995
+ summary_table.add_column("Test Recall", style="red")
996
+
997
+ sorted_results_list = sorted(
998
+ all_experiment_results.items(),
999
+ key=lambda item: (
1000
+ item[1].get("roc_auc", -1) if pd.notna(item[1].get("roc_auc")) else -1
1001
+ ),
1002
+ reverse=True,
1003
+ )
1004
+ best_method_overall_name = "None"
1005
+ best_method_overall_auc = -1.0
1006
+ best_trainable_ml_model_for_api = None
1007
+
1008
+ for method_name_result, metrics_result in sorted_results_list:
1009
+ summary_table.add_row(
1010
+ method_name_result,
1011
+ (
1012
+ f"{metrics_result.get('roc_auc', 'N/A'):.4f}"
1013
+ if pd.notna(metrics_result.get("roc_auc"))
1014
+ else "N/A"
1015
+ ),
1016
+ f"{metrics_result.get('f1_score', 'N/A'):.4f}",
1017
+ f"{metrics_result.get('accuracy', 'N/A'):.4f}",
1018
+ f"{metrics_result.get('precision', 'N/A'):.4f}",
1019
+ f"{metrics_result.get('recall', 'N/A'):.4f}",
1020
+ )
1021
+ current_auc_val_result = metrics_result.get("roc_auc", -1)
1022
+ if (
1023
+ pd.notna(current_auc_val_result)
1024
+ and current_auc_val_result > best_method_overall_auc
1025
+ ):
1026
+ best_method_overall_auc = current_auc_val_result
1027
+ best_method_overall_name = method_name_result
1028
+ if method_name_result in trained_ml_model_objects:
1029
+ best_trainable_ml_model_for_api = trained_ml_model_objects[
1030
+ method_name_result
1031
+ ]
1032
+
1033
+ console.print(summary_table)
1034
+ console.print(
1035
+ f"\n[bold gold1]Best performing method overall for {media_type.upper()} (Test AUC): [white]{best_method_overall_name}[/white] (AUC: {best_method_overall_auc:.4f})[/bold gold1]"
1036
+ )
1037
+
1038
+ results_json_path = os.path.join(
1039
+ experiment_run_output_dir, f"all_experiments_metrics_summary_{media_type}.json"
1040
+ )
1041
+ with open(results_json_path, "w") as f:
1042
+ json.dump(all_experiment_results, f, indent=2, cls=NpEncoder)
1043
+ console.print(
1044
+ f"All experiment metrics summaries for {media_type} saved to [green]{results_json_path}[/green]"
1045
+ )
1046
+
1047
+ plot_roc_curves_all(
1048
+ all_experiment_results,
1049
+ y_meta_test.values,
1050
+ experiment_run_output_dir,
1051
+ media_type,
1052
+ )
1053
+
1054
+ console.print(
1055
+ f"\n[bold]Deployment Artifacts Preparation for {media_type.upper()} (in '{media_type_api_artifacts_subdir}'):[/bold]"
1056
+ )
1057
+
1058
+ joblib.dump(
1059
+ ml_preprocessor.named_steps["imputer"],
1060
+ os.path.join(media_type_api_artifacts_subdir, "deepsafe_meta_imputer.joblib"),
1061
+ )
1062
+ joblib.dump(
1063
+ ml_preprocessor.named_steps["scaler"],
1064
+ os.path.join(media_type_api_artifacts_subdir, "deepsafe_meta_scaler.joblib"),
1065
+ )
1066
+
1067
+ api_feature_cols_path = os.path.join(
1068
+ media_type_api_artifacts_subdir, "deepsafe_meta_feature_columns.json"
1069
+ )
1070
+ if os.path.exists(temp_exp_feature_cols_path):
1071
+ try:
1072
+ with (
1073
+ open(temp_exp_feature_cols_path, "r") as src_f,
1074
+ open(api_feature_cols_path, "w") as dst_f,
1075
+ ):
1076
+ json.dump(json.load(src_f), dst_f, indent=2)
1077
+ console.print(
1078
+ f"Feature columns for {media_type} API saved to [green]{api_feature_cols_path}[/green]"
1079
+ )
1080
+ except Exception as e:
1081
+ console.print(
1082
+ f"[red]Error copying/saving feature columns file: {e}. Manual copy might be needed from {temp_exp_feature_cols_path} to {api_feature_cols_path}[/red]"
1083
+ )
1084
+ else:
1085
+ console.print(
1086
+ f"[yellow]Temporary feature columns file {temp_exp_feature_cols_path} not found. API artifact for feature columns may be missing for {media_type}.[/yellow]"
1087
+ )
1088
+
1089
+ console.print(
1090
+ f"Common imputer, scaler, and feature columns for {media_type} saved for API in '{media_type_api_artifacts_subdir}'."
1091
+ )
1092
+
1093
+ if best_trainable_ml_model_for_api:
1094
+ api_model_joblib_path = os.path.join(
1095
+ media_type_api_artifacts_subdir, "deepsafe_meta_learner.joblib"
1096
+ )
1097
+ joblib.dump(best_trainable_ml_model_for_api, api_model_joblib_path)
1098
+ console.print(
1099
+ f"Best trainable ML meta-learner ([white]{best_method_overall_name}[/white]) for {media_type} saved as '{os.path.basename(api_model_joblib_path)}' in '{media_type_api_artifacts_subdir}'."
1100
+ )
1101
+ console.print(
1102
+ f"The 4 artifacts in '{media_type_api_artifacts_subdir}' are ready for the API."
1103
+ )
1104
+ elif best_method_overall_name.startswith(("Simple", "Provided", "Optimized")):
1105
+ console.print(
1106
+ f"[yellow]The overall best method for {media_type} ([white]{best_method_overall_name}[/white]) is rule-based.[/yellow]"
1107
+ )
1108
+ console.print(
1109
+ f"[yellow]To deploy a trainable ML model, choose the best one from this run and ensure its '.joblib' is saved as 'deepsafe_meta_learner.joblib' inside '{media_type_api_artifacts_subdir}'.[/yellow]"
1110
+ )
1111
+
1112
+ opt_weights_main_dir_path = os.path.join(
1113
+ api_artifacts_dir, f"optimized_grid_average_weights_{media_type}.json"
1114
+ )
1115
+ opt_weights_subdir_path_generic = os.path.join(
1116
+ media_type_api_artifacts_subdir, "optimized_grid_average_weights.json"
1117
+ )
1118
+
1119
+ if "Optimized_Grid_Weighted_Average" in best_method_overall_name:
1120
+ if os.path.exists(opt_weights_main_dir_path):
1121
+ console.print(
1122
+ f" Optimized weights for this method are currently in '{opt_weights_main_dir_path}'. Consider standardizing its location if desired (e.g., to '{opt_weights_subdir_path_generic}')."
1123
+ )
1124
+ elif os.path.exists(
1125
+ opt_weights_subdir_path_generic
1126
+ ): # If you adjust saving logic for weights too
1127
+ console.print(
1128
+ f" Optimized weights for this method are in '{opt_weights_subdir_path_generic}'."
1129
+ )
1130
+ else:
1131
+ console.print(
1132
+ f"[bold red]Error: Could not determine a best trainable model to save for {media_type}. Please review results.[/bold red]"
1133
+ )
1134
+
1135
+ console.rule(
1136
+ f"[bold green]Experimentation Suite for {media_type.upper()} Completed[/bold green]"
1137
+ )
1138
+
1139
+
1140
+ if __name__ == "__main__":
1141
+ parser = argparse.ArgumentParser(
1142
+ description="Run Meta-Learning Experiments for DeepSafe Ensemble."
1143
+ )
1144
+ parser.add_argument(
1145
+ "--media-type",
1146
+ type=str,
1147
+ choices=["image", "video", "audio"],
1148
+ required=True,
1149
+ help="Type of media for which the meta-learner is being trained (image, video, or audio).",
1150
+ )
1151
+ parser.add_argument(
1152
+ "--meta-file",
1153
+ type=str,
1154
+ required=True,
1155
+ help="Path to the media-specific meta-features CSV (e.g., ./meta_data/meta_features_image.csv)",
1156
+ )
1157
+ parser.add_argument(
1158
+ "--output-dir",
1159
+ type=str,
1160
+ default=DEFAULT_EXPERIMENT_OUTPUT_DIR_BASE,
1161
+ help=f"Base directory for saving all experiment-related outputs (default: {DEFAULT_EXPERIMENT_OUTPUT_DIR_BASE}).",
1162
+ )
1163
+ parser.add_argument(
1164
+ "--api-artifacts-dir",
1165
+ type=str,
1166
+ default=DEFAULT_API_ARTIFACTS_DIR,
1167
+ help=f"Directory to save final API-ready artifacts (default: {DEFAULT_API_ARTIFACTS_DIR})",
1168
+ )
1169
+ parser.add_argument(
1170
+ "--optimizer",
1171
+ type=str,
1172
+ choices=["optuna", "gridsearch"],
1173
+ default=OPTIMIZER_CHOICE_DEFAULT,
1174
+ help=f"Hyperparameter optimizer (default: {OPTIMIZER_CHOICE_DEFAULT})",
1175
+ )
1176
+ parser.add_argument(
1177
+ "--optuna-trials",
1178
+ type=int,
1179
+ default=N_OPTUNA_TRIALS_DEFAULT,
1180
+ help=f"Number of Optuna trials (default: {N_OPTUNA_TRIALS_DEFAULT})",
1181
+ )
1182
+ parser.add_argument(
1183
+ "--weights",
1184
+ type=str,
1185
+ default=None,
1186
+ help='JSON string or path to JSON file for custom base model weights (for "Provided_Weighted_Average"). Keys should be base model names (e.g., "npr_deepfakedetection").',
1187
+ )
1188
+
1189
+ args = parser.parse_args()
1190
+
1191
+ if OPTIMIZER_CHOICE_DEFAULT == "optuna" and not OPTIMIZER_AVAILABLE_OPTUNA:
1192
+ console.print(
1193
+ "[yellow]Default optimizer is Optuna, but it's not installed. GridSearchCV will be used if Optuna is chosen via CLI and not available.[/yellow]"
1194
+ )
1195
+ if not XGBOOST_AVAILABLE:
1196
+ console.print(
1197
+ "[yellow]XGBoost not installed. XGBoost experiments will be skipped if its block is reached.[/yellow]"
1198
+ )
1199
+ if not LIGHTGBM_AVAILABLE:
1200
+ console.print(
1201
+ "[yellow]LightGBM not installed. LightGBM experiments will be skipped if its block is reached.[/yellow]"
1202
+ )
1203
+
1204
+ custom_weights_dict_main = None
1205
+ if args.weights:
1206
+ try:
1207
+ if os.path.exists(args.weights):
1208
+ with open(args.weights, "r") as f:
1209
+ custom_weights_dict_main = json.load(f)
1210
+ else:
1211
+ custom_weights_dict_main = json.loads(args.weights)
1212
+ console.print(
1213
+ f"Using provided custom base model weights: {custom_weights_dict_main}"
1214
+ )
1215
+ except Exception as e_weights:
1216
+ console.print(
1217
+ f"[bold red]Error parsing --weights argument: {e_weights}. Proceeding without them.[/bold red]"
1218
+ )
1219
+
1220
+ run_meta_learning_experiments(
1221
+ meta_features_file=args.meta_file,
1222
+ output_dir_base=args.output_dir,
1223
+ api_artifacts_dir=args.api_artifacts_dir,
1224
+ media_type=args.media_type,
1225
+ optimizer_type=args.optimizer,
1226
+ n_optuna_trials_config=args.optuna_trials,
1227
+ provided_custom_weights=custom_weights_dict_main,
1228
+ )
image/aide/.gitignore ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ weights/*.pth
2
+ weights/*.pt
3
+ weights/*/
4
+ hf_cache/
5
+ model_code/
image/aide/Dockerfile ADDED
@@ -0,0 +1,47 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM nvidia/cuda:12.1.1-cudnn8-runtime-ubuntu22.04
2
+
3
+ ENV DEBIAN_FRONTEND=noninteractive
4
+ ENV PYTHONUNBUFFERED=1
5
+
6
+ WORKDIR /app
7
+
8
+ RUN apt-get update && \
9
+ apt-get install -y --no-install-recommends \
10
+ python3 python3-pip python3-dev \
11
+ git wget build-essential && \
12
+ rm -rf /var/lib/apt/lists/*
13
+
14
+ RUN ln -sf /usr/bin/python3 /usr/bin/python
15
+
16
+ RUN pip install --no-cache-dir --upgrade pip "setuptools>=68" wheel
17
+
18
+ # Install PyTorch with CUDA 12.1 (replaces CPU-only arch-conditional install)
19
+ RUN pip install --no-cache-dir \
20
+ torch==2.5.1 torchvision==0.20.1 \
21
+ --index-url https://download.pytorch.org/whl/cu121
22
+
23
+ COPY requirements.txt .
24
+ RUN pip install --no-cache-dir -r requirements.txt
25
+
26
+ RUN git clone https://github.com/shilinyan99/AIDE.git model_code && \
27
+ touch model_code/__init__.py && \
28
+ touch model_code/models/__init__.py && \
29
+ touch model_code/data/__init__.py
30
+
31
+ RUN mkdir -p /app/weights /app/hf_cache
32
+ COPY weights/ /app/weights/
33
+
34
+ COPY app.py .
35
+
36
+ ENV MODEL_PORT=5004
37
+ ENV PRELOAD_MODEL=false
38
+ ENV MODEL_TIMEOUT=600
39
+ ENV AIDE_CHECKPOINT=GenImage_train.pth
40
+ ENV PYTHONPATH=/app/model_code:$PYTHONPATH
41
+
42
+ EXPOSE ${MODEL_PORT}
43
+
44
+ RUN adduser --disabled-password --gecos '' appuser
45
+ USER appuser
46
+
47
+ CMD ["python", "app.py"]
image/aide/app.py ADDED
@@ -0,0 +1,433 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import base64
2
+ import gc
3
+ import io
4
+ import logging
5
+ import os
6
+ import platform
7
+ import sys
8
+ import threading
9
+ import time
10
+ from typing import Any, Dict, Optional
11
+
12
+ import torch
13
+ import uvicorn
14
+ from fastapi import FastAPI, HTTPException
15
+ from fastapi.middleware.cors import CORSMiddleware
16
+ from PIL import Image, ImageFile
17
+ from pydantic import BaseModel
18
+
19
+ ImageFile.LOAD_TRUNCATED_IMAGES = True
20
+
21
+ logging.basicConfig(
22
+ level=logging.INFO,
23
+ format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
24
+ handlers=[logging.StreamHandler(sys.stdout)],
25
+ )
26
+ logger = logging.getLogger(__name__)
27
+
28
+ # ── Path setup ──────────────────────────────────────────────────────────────
29
+ current_dir = os.path.dirname(os.path.abspath(__file__))
30
+ model_code_dir = os.path.join(current_dir, "model_code")
31
+ sys.path.insert(0, model_code_dir)
32
+ sys.path.insert(0, os.path.join(model_code_dir, "models"))
33
+ sys.path.insert(0, os.path.join(model_code_dir, "data"))
34
+
35
+ # ── Compatibility shim ───────────────────────────────────────────────────────
36
+ # AIDE's models/AIDE.py imports `clip` (openai-clip) at module level, but the
37
+ # package is not used during inference β€” only open_clip is. The openai-clip
38
+ # package relies on pkg_resources which was removed in Python 3.13. We inject
39
+ # a lightweight stub so the import succeeds without installing the full package.
40
+ import types as _types
41
+
42
+ if "clip" not in sys.modules:
43
+ _clip_stub = _types.ModuleType("clip")
44
+ sys.modules["clip"] = _clip_stub
45
+
46
+ # ── Config ──────────────────────────────────────────────────────────────────
47
+ MODEL_NAME = "aide_detection"
48
+ WEIGHTS_DIR = os.environ.get("WEIGHTS_DIR", os.path.join(current_dir, "weights"))
49
+ HF_HOME = os.environ.get("HF_HOME", os.path.join(current_dir, "hf_cache"))
50
+ os.environ["HF_HOME"] = HF_HOME
51
+ # The checkpoint is self-contained (includes ConvNeXt weights), so we initialise
52
+ # the architecture with no pretrained weights and load everything from the checkpoint.
53
+ # Set CONVNEXT_PRETRAINED to a HuggingFace tag only if running without a checkpoint.
54
+ CONVNEXT_PRETRAINED = os.environ.get("CONVNEXT_PRETRAINED", None)
55
+
56
+ # Preferred checkpoint filename (GenImage trains on the most diverse generators)
57
+ PREFERRED_CHECKPOINT = os.environ.get("AIDE_CHECKPOINT", "GenImage_train.pth")
58
+
59
+
60
+ def _get_device():
61
+ """Select optimal device: MPS (Apple) > CUDA (NVIDIA) > CPU."""
62
+ override = os.environ.get("DEEPSAFE_DEVICE", "").strip().lower()
63
+ if override == "cpu":
64
+ return torch.device("cpu")
65
+ if override == "cuda" and torch.cuda.is_available():
66
+ return torch.device("cuda")
67
+ if (
68
+ override == "mps"
69
+ and hasattr(torch.backends, "mps")
70
+ and torch.backends.mps.is_available()
71
+ ):
72
+ return torch.device("mps")
73
+ if override:
74
+ pass # Invalid override, fall through to auto-detect
75
+ if (
76
+ platform.system() == "Darwin"
77
+ and hasattr(torch.backends, "mps")
78
+ and torch.backends.mps.is_available()
79
+ ):
80
+ return torch.device("mps")
81
+ if torch.cuda.is_available():
82
+ return torch.device("cuda")
83
+ return torch.device("cpu")
84
+
85
+
86
+ DEVICE = _get_device()
87
+ if DEVICE.type == "cuda":
88
+ torch.backends.cudnn.benchmark = True
89
+ torch.set_float32_matmul_precision("high")
90
+ if DEVICE.type == "cuda":
91
+ logger.info(
92
+ "Device: cuda (%s, %.1f GB VRAM)",
93
+ torch.cuda.get_device_name(0),
94
+ torch.cuda.get_device_properties(0).total_memory / 1024**3,
95
+ )
96
+ else:
97
+ logger.warning(
98
+ "Device: %s (no CUDA available -- check nvidia-container-toolkit)",
99
+ DEVICE,
100
+ )
101
+ PRELOAD_MODEL = os.environ.get("PRELOAD_MODEL", "false").lower() == "true"
102
+ MODEL_TIMEOUT = int(os.environ.get("MODEL_TIMEOUT", "600"))
103
+
104
+ # ── Globals ─────────────────────────────────────────────────────────────────
105
+ model = None
106
+ model_lock = threading.Lock()
107
+ last_used_time = 0
108
+
109
+
110
+ class ImageInput(BaseModel):
111
+ """Request body for /predict."""
112
+
113
+ image_data: str
114
+ threshold: Optional[float] = 0.5
115
+
116
+
117
+ # ── Weight discovery ─────────────────────────────────────────────────────────
118
+
119
+
120
+ def find_aide_checkpoint() -> Optional[str]:
121
+ """
122
+ Return path to the best AIDE checkpoint in WEIGHTS_DIR.
123
+
124
+ Priority order:
125
+ 1. PREFERRED_CHECKPOINT filename (GenImage_train.pth by default)
126
+ 2. Any other .pth file (largest wins)
127
+ """
128
+ if not os.path.exists(WEIGHTS_DIR):
129
+ logger.warning(f"Weights directory not found: {WEIGHTS_DIR}")
130
+ return None
131
+
132
+ # Try the preferred checkpoint first
133
+ preferred = os.path.join(WEIGHTS_DIR, PREFERRED_CHECKPOINT)
134
+ if os.path.exists(preferred):
135
+ logger.info(f"Using preferred checkpoint: {preferred}")
136
+ return preferred
137
+
138
+ # Fall back to the largest available checkpoint
139
+ candidates = [
140
+ os.path.join(WEIGHTS_DIR, f)
141
+ for f in os.listdir(WEIGHTS_DIR)
142
+ if f.endswith(".pth") or f.endswith(".pt")
143
+ ]
144
+ if not candidates:
145
+ logger.warning("No .pth checkpoint found in weights directory.")
146
+ return None
147
+ best = max(candidates, key=os.path.getsize)
148
+ logger.info(f"Using checkpoint: {best}")
149
+ return best
150
+
151
+
152
+ # ── Preprocessing ────────────────────────────────────────────────────────────
153
+
154
+
155
+ def preprocess_image(image_bytes: bytes) -> torch.Tensor:
156
+ """
157
+ Preprocess raw image bytes into AIDE's 5-view tensor.
158
+
159
+ Args:
160
+ image_bytes: Raw bytes of a JPEG/PNG/etc. image.
161
+
162
+ Returns:
163
+ Tensor of shape [1, 5, 3, 256, 256] on CPU.
164
+
165
+ Raises:
166
+ Exception: If bytes cannot be decoded or processed.
167
+ """
168
+ from data.dct import DCT_base_Rec_Module
169
+ from torchvision import transforms
170
+
171
+ pil_image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
172
+
173
+ # Ensure minimum 256x256 so DCT unfold has enough patches
174
+ w, h = pil_image.size
175
+ if w < 256 or h < 256:
176
+ pil_image = pil_image.resize((256, 256), Image.BICUBIC)
177
+
178
+ to_tensor = transforms.ToTensor()
179
+ image_tensor = to_tensor(pil_image) # [3, H, W]
180
+
181
+ # DCT frequency decomposition β†’ 4 patches [3, 32, 32] each
182
+ dct_module = DCT_base_Rec_Module()
183
+ x_minmin, x_maxmax, x_minmin1, x_maxmax1 = dct_module(image_tensor)
184
+
185
+ # Resize all views to 256Γ—256 and normalise with ImageNet stats
186
+ transform = transforms.Compose(
187
+ [
188
+ transforms.Resize([256, 256]),
189
+ transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
190
+ ]
191
+ )
192
+
193
+ x_0 = transform(image_tensor)
194
+ x_minmin = transform(x_minmin)
195
+ x_maxmax = transform(x_maxmax)
196
+ x_minmin1 = transform(x_minmin1)
197
+ x_maxmax1 = transform(x_maxmax1)
198
+
199
+ # Stack β†’ [5, 3, 256, 256], unsqueeze batch β†’ [1, 5, 3, 256, 256]
200
+ stacked = torch.stack([x_minmin, x_maxmax, x_minmin1, x_maxmax1, x_0], dim=0)
201
+ return stacked.unsqueeze(0).to(DEVICE)
202
+
203
+
204
+ # ── Model loading ─────────────────────────────────────────────────────────────
205
+
206
+
207
+ def load_model_internal():
208
+ """Load AIDE_Model onto CPU with the best available checkpoint."""
209
+ global model, last_used_time
210
+
211
+ with model_lock:
212
+ if model is not None:
213
+ last_used_time = time.time()
214
+ return
215
+
216
+ logger.info("Loading AIDE model...")
217
+ try:
218
+ import models.AIDE as AIDE_module
219
+
220
+ aide_model = AIDE_module.AIDE(
221
+ resnet_path=None,
222
+ convnext_path=CONVNEXT_PRETRAINED,
223
+ )
224
+ aide_model.to(DEVICE)
225
+
226
+ checkpoint_path = find_aide_checkpoint()
227
+ if checkpoint_path:
228
+ logger.info(f"Loading checkpoint: {checkpoint_path}")
229
+ ckpt = torch.load(checkpoint_path, map_location=DEVICE)
230
+ if isinstance(ckpt, dict):
231
+ state_dict = ckpt.get("model") or ckpt.get("state_dict") or ckpt
232
+ else:
233
+ state_dict = ckpt
234
+ # Strip DataParallel "module." prefix if present
235
+ cleaned = {
236
+ k[7:] if k.startswith("module.") else k: v
237
+ for k, v in state_dict.items()
238
+ }
239
+ missing, unexpected = aide_model.load_state_dict(cleaned, strict=False)
240
+ logger.info(
241
+ f"Checkpoint loaded. Missing keys: {len(missing)}, "
242
+ f"Unexpected keys: {len(unexpected)}"
243
+ )
244
+ else:
245
+ logger.warning(
246
+ "No checkpoint found β€” model uses pretrained-only weights. "
247
+ "Run download_weights.sh to fetch the fine-tuned checkpoint."
248
+ )
249
+
250
+ # Switch to inference mode (no gradient tracking, batch-norm uses running stats)
251
+ aide_model.train(mode=False)
252
+ model = aide_model
253
+ last_used_time = time.time()
254
+ logger.info("AIDE model ready.")
255
+
256
+ except Exception as exc:
257
+ logger.exception(f"Failed to load AIDE model: {exc}")
258
+ model = None
259
+ raise
260
+ finally:
261
+ gc.collect()
262
+
263
+
264
+ def ensure_model_loaded():
265
+ """Load model on first request (lazy loading)."""
266
+ global last_used_time
267
+ if model is None:
268
+ load_model_internal()
269
+ else:
270
+ last_used_time = time.time()
271
+
272
+
273
+ def unload_model_if_idle():
274
+ """Evict model from RAM after MODEL_TIMEOUT seconds of inactivity."""
275
+ global model
276
+ if model is None or PRELOAD_MODEL:
277
+ return
278
+ with model_lock:
279
+ if model is not None and (time.time() - last_used_time > MODEL_TIMEOUT):
280
+ logger.info("Unloading idle AIDE model to free RAM.")
281
+ del model
282
+ model = None
283
+ gc.collect()
284
+
285
+
286
+ # ── FastAPI app ───────────────────────────────────────────────────────────────
287
+
288
+ app = FastAPI(
289
+ title="AIDE Deepfake Detection Service",
290
+ description=(
291
+ "AI-generated image detection using AIDE (ICLR 2025) β€” "
292
+ "hybrid DCT frequency analysis + ConvNeXt-xxlarge semantic features."
293
+ ),
294
+ version="1.0.0",
295
+ )
296
+
297
+ app.add_middleware(
298
+ CORSMiddleware,
299
+ allow_origins=["*"],
300
+ allow_credentials=True,
301
+ allow_methods=["*"],
302
+ allow_headers=["*"],
303
+ )
304
+
305
+
306
+ @app.get("/")
307
+ async def root():
308
+ """Root endpoint with service information."""
309
+ return {
310
+ "model_name": MODEL_NAME,
311
+ "description": "AIDE (ICLR 2025) AI-generated image detector",
312
+ "device": str(DEVICE),
313
+ "model_loaded": model is not None,
314
+ }
315
+
316
+
317
+ def _gpu_health_info() -> dict:
318
+ """Return GPU metrics for the health endpoint."""
319
+ if torch.cuda.is_available() and DEVICE.type == "cuda":
320
+ return {
321
+ "gpu_name": torch.cuda.get_device_name(0),
322
+ "vram_used_mb": round(torch.cuda.memory_allocated(0) / 1024**2),
323
+ "vram_total_mb": round(
324
+ torch.cuda.get_device_properties(0).total_memory / 1024**2
325
+ ),
326
+ }
327
+ return {}
328
+
329
+
330
+ @app.get("/health")
331
+ async def health():
332
+ """Health check endpoint."""
333
+ return {
334
+ "status": "healthy",
335
+ "model_name": MODEL_NAME,
336
+ "device": str(DEVICE),
337
+ "model_loaded": model is not None,
338
+ **_gpu_health_info(),
339
+ }
340
+
341
+
342
+ @app.post("/unload")
343
+ async def unload_model_endpoint():
344
+ """Manually unload the model to free RAM."""
345
+ global model
346
+ if model is None:
347
+ return {"status": "not_loaded"}
348
+ del model
349
+ model = None
350
+ gc.collect()
351
+ return {"status": "success", "message": "Model unloaded."}
352
+
353
+
354
+ @app.post("/predict")
355
+ async def predict(image_input: ImageInput) -> Dict[str, Any]:
356
+ """
357
+ Predict whether the submitted image is AI-generated.
358
+
359
+ Args:
360
+ image_input: Base64-encoded image and optional classification threshold.
361
+
362
+ Returns:
363
+ Dict with model name, fake probability, binary prediction, class label, and inference time.
364
+ """
365
+ try:
366
+ ensure_model_loaded()
367
+ if model is None:
368
+ raise HTTPException(status_code=503, detail="Model not loaded.")
369
+
370
+ start = time.time()
371
+
372
+ try:
373
+ image_bytes = base64.b64decode(image_input.image_data)
374
+ input_tensor = preprocess_image(image_bytes)
375
+ except Exception as exc:
376
+ raise HTTPException(status_code=400, detail=f"Invalid image data: {exc}")
377
+
378
+ with torch.no_grad():
379
+ logits = model(input_tensor) # [1, 2]
380
+ probs = torch.softmax(logits, dim=-1) # [1, 2]
381
+ probability_fake = probs[0, 1].item()
382
+
383
+ prediction = 1 if probability_fake >= image_input.threshold else 0
384
+ class_label = "fake" if prediction == 1 else "real"
385
+ inference_time = time.time() - start
386
+
387
+ logger.info(
388
+ f"Prediction: {class_label} (prob={probability_fake:.4f}, {inference_time:.3f}s)"
389
+ )
390
+
391
+ if not PRELOAD_MODEL and MODEL_TIMEOUT > 0:
392
+ threading.Timer(MODEL_TIMEOUT + 5.0, unload_model_if_idle).start()
393
+
394
+ return {
395
+ "model": MODEL_NAME,
396
+ "probability": float(probability_fake),
397
+ "prediction": int(prediction),
398
+ "class": class_label,
399
+ "inference_time": float(inference_time),
400
+ }
401
+
402
+ except HTTPException:
403
+ raise
404
+ except Exception as exc:
405
+ logger.exception(f"Prediction error: {exc}")
406
+ raise HTTPException(status_code=500, detail=str(exc))
407
+
408
+
409
+ @app.on_event("startup")
410
+ async def startup_event():
411
+ """Startup handler β€” preloads model if PRELOAD_MODEL=true, else lazy-loads."""
412
+ if PRELOAD_MODEL:
413
+ logger.info("Preloading AIDE model at startup.")
414
+ try:
415
+ load_model_internal()
416
+ except Exception as exc:
417
+ logger.error(f"Preload failed: {exc}")
418
+ else:
419
+ logger.info("AIDE service ready β€” model loads on first request.")
420
+
421
+ if not PRELOAD_MODEL and MODEL_TIMEOUT > 0:
422
+
423
+ def _periodic_check():
424
+ unload_model_if_idle()
425
+ threading.Timer(MODEL_TIMEOUT / 2.0, _periodic_check).start()
426
+
427
+ threading.Timer(MODEL_TIMEOUT / 2.0, _periodic_check).start()
428
+
429
+
430
+ if __name__ == "__main__":
431
+ port = int(os.environ.get("MODEL_PORT", 5004))
432
+ logger.info(f"Starting AIDE service on port {port}")
433
+ uvicorn.run(app, host="0.0.0.0", port=port)