Instructions to use deepsafe/deepsafe-services with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use deepsafe/deepsafe-services with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("deepsafe/deepsafe-services", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
sync from GitHub (0154d02)
Browse filesThis view is limited to 50 files because it contains too many changes. Β See raw diff
- audio/aasist3/.gitignore +1 -0
- audio/aasist3/Dockerfile +35 -0
- audio/aasist3/api.py +345 -0
- audio/aasist3/model/__init__.py +1 -0
- audio/aasist3/model/branch.py +34 -0
- audio/aasist3/model/full_model.py +139 -0
- audio/aasist3/model/gat.py +99 -0
- audio/aasist3/model/hs_gal.py +176 -0
- audio/aasist3/model/kan.py +213 -0
- audio/aasist3/model/pool.py +45 -0
- audio/aasist3/model/residual.py +56 -0
- audio/aasist3/model/wav2vec.py +82 -0
- audio/aasist3/requirements.txt +11 -0
- audio/nes2net/.gitignore +1 -0
- audio/nes2net/Dockerfile +41 -0
- audio/nes2net/api.py +315 -0
- audio/nes2net/model_scripts/__init__.py +0 -0
- audio/nes2net/model_scripts/wav2vec2_Nes2Net_X.py +317 -0
- audio/nes2net/requirements.txt +13 -0
- audio/safeear/Dockerfile +49 -0
- audio/safeear/api.py +321 -0
- audio/safeear/download_weights.sh +23 -0
- audio/safeear/requirements.txt +15 -0
- audio/shiftyspeech/Dockerfile +43 -0
- audio/shiftyspeech/api.py +315 -0
- audio/shiftyspeech/evaluate.py +210 -0
- audio/shiftyspeech/requirements.txt +14 -0
- audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/.env +2 -0
- audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/LICENSE +21 -0
- audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/RawBoost.py +143 -0
- audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/Simplified_CM_solution.py +227 -0
- audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/data_utils.py +292 -0
- audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/model.py +603 -0
- audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/requirements.txt +4 -0
- audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/startup_config.py +60 -0
- audio/shiftyspeech/synthetic_speech_detection/SSL_Anti-spoofing/train.py +446 -0
- audio/shiftyspeech/tests/__init__.py +0 -0
- audio/shiftyspeech/tests/test_api.py +253 -0
- audio/sonics/Dockerfile +36 -0
- audio/sonics/app.py +285 -0
- audio/sonics/requirements.txt +22 -0
- ensemble-core/Dockerfile +10 -0
- ensemble-core/main.py +23 -0
- ensemble-core/requirements.txt +1 -0
- ensemble-core/scripts/create_dataset.py +169 -0
- ensemble-core/scripts/meta_feature_generator.py +366 -0
- ensemble-core/scripts/train_meta_learner_advanced.py +1228 -0
- image/aide/.gitignore +5 -0
- image/aide/Dockerfile +47 -0
- 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)
|