deepsafe commited on
Commit
505afd3
·
verified ·
1 Parent(s): 442dad0

Upload DeepSafe services: model weights and inference code

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +94 -0
  2. audio/aasist3/.gitignore +1 -0
  3. audio/aasist3/Dockerfile +33 -0
  4. audio/aasist3/api.py +314 -0
  5. audio/aasist3/model/__init__.py +1 -0
  6. audio/aasist3/model/branch.py +34 -0
  7. audio/aasist3/model/full_model.py +139 -0
  8. audio/aasist3/model/gat.py +99 -0
  9. audio/aasist3/model/hs_gal.py +176 -0
  10. audio/aasist3/model/kan.py +213 -0
  11. audio/aasist3/model/pool.py +45 -0
  12. audio/aasist3/model/residual.py +56 -0
  13. audio/aasist3/model/wav2vec.py +82 -0
  14. audio/aasist3/requirements.txt +11 -0
  15. audio/aasist3/weights/config.json +45 -0
  16. audio/aasist3/weights/model.safetensors +3 -0
  17. audio/nes2net/.gitignore +1 -0
  18. audio/nes2net/Dockerfile +45 -0
  19. audio/nes2net/api.py +284 -0
  20. audio/nes2net/model_scripts/__init__.py +0 -0
  21. audio/nes2net/model_scripts/wav2vec2_Nes2Net_X.py +317 -0
  22. audio/nes2net/requirements.txt +13 -0
  23. audio/nes2net/weights/nes2net_itw_valaug.pt +3 -0
  24. audio/nes2net/weights/xlsr2_300m.pt +3 -0
  25. audio/safeear/Dockerfile +39 -0
  26. audio/safeear/api.py +290 -0
  27. audio/safeear/download_weights.sh +23 -0
  28. audio/safeear/requirements.txt +15 -0
  29. audio/safeear/safeear_repo/.gitignore +167 -0
  30. audio/safeear/safeear_repo/LICENSE +23 -0
  31. audio/safeear/safeear_repo/README.md +134 -0
  32. audio/safeear/safeear_repo/assert/ASVSpoof-results.png +3 -0
  33. audio/safeear/safeear_repo/assert/Fig1.jpg +3 -0
  34. audio/safeear/safeear_repo/assert/SafeEar_logo.jpg +3 -0
  35. audio/safeear/safeear_repo/assert/exp1.png +3 -0
  36. audio/safeear/safeear_repo/assert/overall.gif +3 -0
  37. audio/safeear/safeear_repo/assert/overall.mp4 +3 -0
  38. audio/safeear/safeear_repo/assert/overall.png +3 -0
  39. audio/safeear/safeear_repo/assert/safe-space.png +0 -0
  40. audio/safeear/safeear_repo/config/train19.yaml +87 -0
  41. audio/safeear/safeear_repo/config/train21.yaml +87 -0
  42. audio/safeear/safeear_repo/datas/ASVSpoof2019/ASVspoof2019.LA.cm.dev.trl.txt +0 -0
  43. audio/safeear/safeear_repo/datas/ASVSpoof2019/ASVspoof2019.LA.cm.eval.trl.txt +0 -0
  44. audio/safeear/safeear_repo/datas/ASVSpoof2019/ASVspoof2019.LA.cm.train.trn.txt +0 -0
  45. audio/safeear/safeear_repo/datas/ASVSpoof2019/dev.tsv +0 -0
  46. audio/safeear/safeear_repo/datas/ASVSpoof2019/eval.tsv +0 -0
  47. audio/safeear/safeear_repo/datas/ASVSpoof2019/train.tsv +0 -0
  48. audio/safeear/safeear_repo/datas/ASVSpoof2021/ASVspoof2021.LA.cm.eval.trl.txt +0 -0
  49. audio/safeear/safeear_repo/datas/ASVSpoof2021/eval.tsv +0 -0
  50. audio/safeear/safeear_repo/datas/dump_hubert_avg_feature.py +109 -0
.gitattributes CHANGED
@@ -33,3 +33,97 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ audio/safeear/safeear_repo/assert/ASVSpoof-results.png filter=lfs diff=lfs merge=lfs -text
37
+ audio/safeear/safeear_repo/assert/Fig1.jpg filter=lfs diff=lfs merge=lfs -text
38
+ audio/safeear/safeear_repo/assert/SafeEar_logo.jpg filter=lfs diff=lfs merge=lfs -text
39
+ audio/safeear/safeear_repo/assert/exp1.png filter=lfs diff=lfs merge=lfs -text
40
+ audio/safeear/safeear_repo/assert/overall.gif filter=lfs diff=lfs merge=lfs -text
41
+ audio/safeear/safeear_repo/assert/overall.mp4 filter=lfs diff=lfs merge=lfs -text
42
+ audio/safeear/safeear_repo/assert/overall.png filter=lfs diff=lfs merge=lfs -text
43
+ audio/safeear/safeear_repo/fairseq_ours/docs/fairseq.gif filter=lfs diff=lfs merge=lfs -text
44
+ audio/safeear/safeear_repo/fairseq_ours/examples/MMPT/videoclip.png filter=lfs diff=lfs merge=lfs -text
45
+ audio/safeear/safeear_repo/fairseq_ours/examples/MMPT/vlm.png filter=lfs diff=lfs merge=lfs -text
46
+ audio/safeear/safeear_repo/fairseq_ours/examples/hubert/tests/6313-76958-0021.flac filter=lfs diff=lfs merge=lfs -text
47
+ audio/safeear/safeear_repo/fairseq_ours/examples/textless_nlp/speech-resynth/img/fig.png filter=lfs diff=lfs merge=lfs -text
48
+ audio/shiftyspeech/Images/logo-png.png filter=lfs diff=lfs merge=lfs -text
49
+ audio/shiftyspeech/Images/logo-transparent-png.png filter=lfs diff=lfs merge=lfs -text
50
+ audio/shiftyspeech/Images/shiftyspeech_logo.png filter=lfs diff=lfs merge=lfs -text
51
+ audio/shiftyspeech/speech_synthesis/tts/Grad-TTS/out/sample_1.wav filter=lfs diff=lfs merge=lfs -text
52
+ audio/shiftyspeech/speech_synthesis/tts/Grad-TTS/out/sample_2.wav filter=lfs diff=lfs merge=lfs -text
53
+ audio/shiftyspeech/speech_synthesis/tts/Grad-TTS/resources/reverse-diffusion.gif filter=lfs diff=lfs merge=lfs -text
54
+ audio/shiftyspeech/speech_synthesis/vocoders/FreeV/figure/compare_table.png filter=lfs diff=lfs merge=lfs -text
55
+ audio/shiftyspeech/speech_synthesis/vocoders/FreeV/figure/overall.png filter=lfs diff=lfs merge=lfs -text
56
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvgan/BigVGAN/demo/examples/dance_24k.wav filter=lfs diff=lfs merge=lfs -text
57
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvgan/BigVGAN/demo/examples/hifitts_44k.wav filter=lfs diff=lfs merge=lfs -text
58
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvgan/BigVGAN/demo/examples/jensen_24k.wav filter=lfs diff=lfs merge=lfs -text
59
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvgan/BigVGAN/demo/examples/libritts_24k.wav filter=lfs diff=lfs merge=lfs -text
60
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvgan/BigVGAN/demo/examples/megalovania_24k.wav filter=lfs diff=lfs merge=lfs -text
61
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvgan/BigVGAN/demo/examples/musdbhq_44k.wav filter=lfs diff=lfs merge=lfs -text
62
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvgan/BigVGAN/demo/examples/musiccaps1_44k.wav filter=lfs diff=lfs merge=lfs -text
63
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvgan/BigVGAN/demo/examples/musiccaps2_44k.wav filter=lfs diff=lfs merge=lfs -text
64
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvgan/BigVGAN/demo/examples/queen_24k.wav filter=lfs diff=lfs merge=lfs -text
65
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvgan/BigVGAN/filelists/LibriTTS/train-full.txt filter=lfs diff=lfs merge=lfs -text
66
+ audio/shiftyspeech/speech_synthesis/vocoders/bigvsan/LibriTTS/train-full.txt filter=lfs diff=lfs merge=lfs -text
67
+ audio/shiftyspeech/speech_synthesis/vocoders/iSTFTNet-pytorch/iSTFTnet.PNG filter=lfs diff=lfs merge=lfs -text
68
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/datasets/metadata/libritts_train_clean_360_audiopath_text_sid_train.txt filter=lfs diff=lfs merge=lfs -text
69
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/loss.png filter=lfs diff=lfs merge=lfs -text
70
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/model_architecture.png filter=lfs diff=lfs merge=lfs -text
71
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c16/2004_147967_000029_000002.wav filter=lfs diff=lfs merge=lfs -text
72
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c16/337_126286_000008_000000.wav filter=lfs diff=lfs merge=lfs -text
73
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c16/3537_5704_000008_000005.wav filter=lfs diff=lfs merge=lfs -text
74
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c16/5319_84357_000005_000004.wav filter=lfs diff=lfs merge=lfs -text
75
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c16/6294_86679_000035_000004.wav filter=lfs diff=lfs merge=lfs -text
76
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c16/949_134657_000002_000005.wav filter=lfs diff=lfs merge=lfs -text
77
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c32/2004_147967_000029_000002.wav filter=lfs diff=lfs merge=lfs -text
78
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c32/337_126286_000008_000000.wav filter=lfs diff=lfs merge=lfs -text
79
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c32/3537_5704_000008_000005.wav filter=lfs diff=lfs merge=lfs -text
80
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c32/5319_84357_000005_000004.wav filter=lfs diff=lfs merge=lfs -text
81
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c32/6294_86679_000035_000004.wav filter=lfs diff=lfs merge=lfs -text
82
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/c32/949_134657_000002_000005.wav filter=lfs diff=lfs merge=lfs -text
83
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/ground_truth/2004_147967_000029_000002.wav filter=lfs diff=lfs merge=lfs -text
84
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/ground_truth/337_126286_000008_000000.wav filter=lfs diff=lfs merge=lfs -text
85
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/ground_truth/3537_5704_000008_000005.wav filter=lfs diff=lfs merge=lfs -text
86
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/ground_truth/5319_84357_000005_000004.wav filter=lfs diff=lfs merge=lfs -text
87
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/ground_truth/6294_86679_000035_000004.wav filter=lfs diff=lfs merge=lfs -text
88
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/ground_truth/949_134657_000002_000005.wav filter=lfs diff=lfs merge=lfs -text
89
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/official_c16/2004_147967_000029_000002.wav filter=lfs diff=lfs merge=lfs -text
90
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/official_c16/337_126286_000008_000000.wav filter=lfs diff=lfs merge=lfs -text
91
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/official_c16/3537_5704_000008_000005.wav filter=lfs diff=lfs merge=lfs -text
92
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/official_c16/5319_84357_000005_000004.wav filter=lfs diff=lfs merge=lfs -text
93
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/official_c16/6294_86679_000035_000004.wav filter=lfs diff=lfs merge=lfs -text
94
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/official_c32/2004_147967_000029_000002.wav filter=lfs diff=lfs merge=lfs -text
95
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/official_c32/337_126286_000008_000000.wav filter=lfs diff=lfs merge=lfs -text
96
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/official_c32/3537_5704_000008_000005.wav filter=lfs diff=lfs merge=lfs -text
97
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/official_c32/5319_84357_000005_000004.wav filter=lfs diff=lfs merge=lfs -text
98
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/seen/official_c32/6294_86679_000035_000004.wav filter=lfs diff=lfs merge=lfs -text
99
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c16/1089_134686_000007_000005.wav filter=lfs diff=lfs merge=lfs -text
100
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c16/3575_170457_000037_000002.wav filter=lfs diff=lfs merge=lfs -text
101
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c16/4507_16021_000029_000005.wav filter=lfs diff=lfs merge=lfs -text
102
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c16/7021_85628_000037_000000.wav filter=lfs diff=lfs merge=lfs -text
103
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c16/7176_92135_000006_000005.wav filter=lfs diff=lfs merge=lfs -text
104
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c16/8224_274384_000016_000000.wav filter=lfs diff=lfs merge=lfs -text
105
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c32/1089_134686_000007_000005.wav filter=lfs diff=lfs merge=lfs -text
106
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c32/3575_170457_000037_000002.wav filter=lfs diff=lfs merge=lfs -text
107
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c32/4507_16021_000029_000005.wav filter=lfs diff=lfs merge=lfs -text
108
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c32/7021_85628_000037_000000.wav filter=lfs diff=lfs merge=lfs -text
109
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c32/7176_92135_000006_000005.wav filter=lfs diff=lfs merge=lfs -text
110
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/c32/8224_274384_000016_000000.wav filter=lfs diff=lfs merge=lfs -text
111
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/ground_truth/1089_134686_000007_000005.wav filter=lfs diff=lfs merge=lfs -text
112
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/ground_truth/3575_170457_000037_000002.wav filter=lfs diff=lfs merge=lfs -text
113
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/ground_truth/4507_16021_000029_000005.wav filter=lfs diff=lfs merge=lfs -text
114
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/ground_truth/7021_85628_000037_000000.wav filter=lfs diff=lfs merge=lfs -text
115
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/ground_truth/7176_92135_000006_000005.wav filter=lfs diff=lfs merge=lfs -text
116
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/ground_truth/8224_274384_000016_000000.wav filter=lfs diff=lfs merge=lfs -text
117
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/official_c16/1089_134686_000007_000005.wav filter=lfs diff=lfs merge=lfs -text
118
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/official_c16/3575_170457_000037_000002.wav filter=lfs diff=lfs merge=lfs -text
119
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/official_c16/7021_85628_000037_000000.wav filter=lfs diff=lfs merge=lfs -text
120
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/official_c16/7176_92135_000006_000005.wav filter=lfs diff=lfs merge=lfs -text
121
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/official_c16/8224_274384_000016_000000.wav filter=lfs diff=lfs merge=lfs -text
122
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/official_c32/1089_134686_000007_000005.wav filter=lfs diff=lfs merge=lfs -text
123
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/official_c32/3575_170457_000037_000002.wav filter=lfs diff=lfs merge=lfs -text
124
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/official_c32/7021_85628_000037_000000.wav filter=lfs diff=lfs merge=lfs -text
125
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/official_c32/7176_92135_000006_000005.wav filter=lfs diff=lfs merge=lfs -text
126
+ audio/shiftyspeech/speech_synthesis/vocoders/univnet/docs/samples/unseen/official_c32/8224_274384_000016_000000.wav filter=lfs diff=lfs merge=lfs -text
127
+ image/aide/model_code/docs/Chameleon.jpg filter=lfs diff=lfs merge=lfs -text
128
+ image/aide/model_code/docs/network.png filter=lfs diff=lfs merge=lfs -text
129
+ video/fake-stormer/model_code/demo/method.png filter=lfs diff=lfs merge=lfs -text
audio/aasist3/.gitignore ADDED
@@ -0,0 +1 @@
 
 
1
+ weights/
audio/aasist3/Dockerfile ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM python:3.10-slim
2
+
3
+ # Install system dependencies (no git/build-essential needed)
4
+ RUN apt-get update && apt-get install -y \
5
+ ffmpeg \
6
+ libsndfile1 \
7
+ && rm -rf /var/lib/apt/lists/*
8
+
9
+ WORKDIR /app
10
+
11
+ # Install Python dependencies
12
+ COPY requirements.txt .
13
+ RUN pip install --no-cache-dir -r requirements.txt
14
+
15
+ # Pre-cache wav2vec2 config (not weights, just the config JSON ~1KB)
16
+ RUN python -c "from transformers import Wav2Vec2Config; Wav2Vec2Config.from_pretrained('facebook/wav2vec2-large-xlsr-53', cache_dir='/app/w2v_cache')"
17
+
18
+ # Copy model code and API
19
+ COPY model/ /app/model/
20
+ COPY api.py .
21
+
22
+ # Expected weight files in /app/weights/:
23
+ # - model.safetensors (AASIST3 checkpoint)
24
+ # - config.json (model configuration)
25
+ # These are mounted from the host via docker-compose volumes.
26
+
27
+ EXPOSE 8005
28
+
29
+ # Drop root privileges
30
+ RUN adduser --disabled-password --gecos '' appuser
31
+ USER appuser
32
+
33
+ CMD ["python", "api.py"]
audio/aasist3/api.py ADDED
@@ -0,0 +1,314 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ SAMPLE_RATE = 16000
79
+ TARGET_SAMPLES = 64600 # ~4.04 seconds at 16kHz
80
+
81
+ # Global model instance
82
+ model = None
83
+
84
+
85
+ class AudioInput(BaseModel):
86
+ """Request schema for audio deepfake detection."""
87
+
88
+ audio_data: str = Field(
89
+ ..., description="Base64 encoded audio string (WAV/MP3/etc)"
90
+ )
91
+ threshold: Optional[float] = Field(
92
+ 0.5, ge=0.0, le=1.0, description="Classification threshold"
93
+ )
94
+
95
+
96
+ app = FastAPI(
97
+ title="AASIST3 Audio Deepfake Detection API",
98
+ description=(
99
+ "Service for detecting synthetic speech using the "
100
+ "AASIST3 model (HuggingFace wav2vec 2.0 + AASIST "
101
+ "with KAN layers)."
102
+ ),
103
+ version="1.0.0",
104
+ )
105
+
106
+
107
+ def load_model():
108
+ """Load the AASIST3 model from pretrained weights.
109
+
110
+ Returns:
111
+ The loaded model, or None if loading fails.
112
+ """
113
+ global model
114
+ if model is not None:
115
+ return model
116
+
117
+ logger.info(f"Loading AASIST3 model onto {DEVICE}...")
118
+
119
+ if AASIST3Model is None:
120
+ logger.error("AASIST3 model class not available.")
121
+ return None
122
+
123
+ weights_safetensors = os.path.join(WEIGHTS_DIR, "model.safetensors")
124
+ weights_config = os.path.join(WEIGHTS_DIR, "config.json")
125
+
126
+ if not os.path.exists(weights_safetensors):
127
+ logger.error(f"Model weights not found at {weights_safetensors}")
128
+ return None
129
+
130
+ if not os.path.exists(weights_config):
131
+ logger.error(f"Model config not found at {weights_config}")
132
+ return None
133
+
134
+ try:
135
+ model = AASIST3Model.from_pretrained(WEIGHTS_DIR)
136
+ model.to(DEVICE)
137
+ # Set model to inference mode (disables dropout, batchnorm)
138
+ model.train(False)
139
+
140
+ logger.info("AASIST3 model loaded successfully.")
141
+ return model
142
+ except Exception as e:
143
+ logger.exception(f"Failed to load AASIST3 model: {e}")
144
+ model = None
145
+ return None
146
+
147
+
148
+ @app.on_event("startup")
149
+ async def startup_event():
150
+ """Load model on service startup."""
151
+ load_model()
152
+
153
+
154
+ @app.get("/health")
155
+ async def health():
156
+ """Health check endpoint."""
157
+ return {
158
+ "status": "healthy" if model is not None else "degraded",
159
+ "model": MODEL_NAME,
160
+ "model_id": MODEL_ID,
161
+ "device": str(DEVICE),
162
+ "weights_found": os.path.exists(os.path.join(WEIGHTS_DIR, "model.safetensors")),
163
+ }
164
+
165
+
166
+ def _load_audio_bytes(audio_bytes: bytes) -> tuple:
167
+ """Load audio from raw bytes using soundfile with torchaudio fallback.
168
+
169
+ Args:
170
+ audio_bytes: Raw audio file bytes.
171
+
172
+ Returns:
173
+ Tuple of (audio_numpy_array, sample_rate).
174
+
175
+ Raises:
176
+ ValueError: If audio cannot be loaded by any backend.
177
+ """
178
+ # Try soundfile first (handles WAV, FLAC natively)
179
+ sf_error = None
180
+ try:
181
+ audio, sr = sf.read(io.BytesIO(audio_bytes), dtype="float32")
182
+ if audio.ndim > 1:
183
+ audio = audio.mean(axis=1) # Convert to mono
184
+ return audio, sr
185
+ except Exception as sf_err:
186
+ sf_error = sf_err
187
+ logger.debug(f"soundfile failed, trying torchaudio: {sf_err}")
188
+
189
+ # Fallback to torchaudio (handles MP3, compressed formats)
190
+ try:
191
+ import torchaudio
192
+
193
+ buf = io.BytesIO(audio_bytes)
194
+ waveform, sr = torchaudio.load(buf)
195
+ if waveform.shape[0] > 1:
196
+ waveform = waveform.mean(dim=0, keepdim=True)
197
+ return waveform.squeeze(0).numpy(), sr
198
+ except Exception as ta_err:
199
+ raise ValueError(
200
+ f"Failed to load audio with soundfile and torchaudio: "
201
+ f"sf={sf_error}, ta={ta_err}"
202
+ )
203
+
204
+
205
+ def preprocess_audio(audio_bytes: bytes) -> torch.Tensor:
206
+ """Preprocess audio for AASIST3 inference.
207
+
208
+ Loads audio, resamples to 16kHz mono, and zero-pads or
209
+ truncates to TARGET_SAMPLES.
210
+
211
+ Args:
212
+ audio_bytes: Raw audio file bytes.
213
+
214
+ Returns:
215
+ Audio tensor of shape (1, TARGET_SAMPLES).
216
+
217
+ Raises:
218
+ ValueError: If audio preprocessing fails.
219
+ """
220
+ try:
221
+ logger.info("Starting audio preprocessing...")
222
+ audio, sr = _load_audio_bytes(audio_bytes)
223
+ logger.info(f"Audio loaded. Length: {len(audio)} samples at {sr}Hz")
224
+
225
+ # Resample to 16kHz if needed
226
+ if sr != SAMPLE_RATE:
227
+ import torchaudio
228
+
229
+ resampler = torchaudio.transforms.Resample(
230
+ orig_freq=sr, new_freq=SAMPLE_RATE
231
+ )
232
+ audio_tensor = torch.FloatTensor(audio).unsqueeze(0)
233
+ audio_tensor = resampler(audio_tensor).squeeze(0)
234
+ audio = audio_tensor.numpy()
235
+ logger.info(
236
+ f"Resampled from {sr}Hz to {SAMPLE_RATE}Hz. "
237
+ f"New length: {len(audio)} samples"
238
+ )
239
+
240
+ # Zero-pad or truncate to TARGET_SAMPLES
241
+ if len(audio) >= TARGET_SAMPLES:
242
+ audio = audio[:TARGET_SAMPLES]
243
+ else:
244
+ pad_length = TARGET_SAMPLES - len(audio)
245
+ audio = np.pad(audio, (0, pad_length), mode="constant")
246
+
247
+ logger.info(f"Audio padded/trimmed to {TARGET_SAMPLES} samples")
248
+
249
+ audio_tensor = torch.FloatTensor(audio).unsqueeze(0).to(DEVICE)
250
+ return audio_tensor
251
+ except Exception as e:
252
+ logger.error(f"Error preprocessing audio: {e}")
253
+ raise ValueError(f"Audio preprocessing failed: {str(e)}")
254
+
255
+
256
+ @app.post("/predict")
257
+ async def predict(input_data: AudioInput):
258
+ """Run deepfake detection on base64-encoded audio.
259
+
260
+ The model outputs 2 logits: [bonafide_score, spoof_score].
261
+ Class 0 = bonafide (real), Class 1 = spoof (fake).
262
+ The returned probability is the spoof/fake probability.
263
+ """
264
+ if model is None:
265
+ if load_model() is None:
266
+ raise HTTPException(status_code=503, detail="Model not loaded")
267
+
268
+ try:
269
+ start_time = time.time()
270
+ logger.info(
271
+ f"Prediction request. Data size: " f"{len(input_data.audio_data)} chars"
272
+ )
273
+
274
+ # Decode base64 audio
275
+ audio_bytes = base64.b64decode(input_data.audio_data)
276
+
277
+ # Preprocess
278
+ audio_tensor = preprocess_audio(audio_bytes)
279
+
280
+ # Inference
281
+ logger.info("Starting model inference...")
282
+ with torch.no_grad():
283
+ output = model(audio_tensor)
284
+
285
+ # output shape: [batch, 2]
286
+ # Index 0 = bonafide logit, Index 1 = spoof logit
287
+ probs = torch.softmax(output, dim=1)
288
+ prob_fake = probs[0, 1].item()
289
+
290
+ prediction = 1 if prob_fake >= input_data.threshold else 0
291
+ verdict = "fake" if prediction == 1 else "real"
292
+ inference_time = time.time() - start_time
293
+
294
+ logger.info(
295
+ f"Prediction: {verdict} (prob_fake={prob_fake:.4f}, "
296
+ f"time={inference_time:.3f}s)"
297
+ )
298
+
299
+ return {
300
+ "model": MODEL_NAME,
301
+ "probability": float(prob_fake),
302
+ "prediction": int(prediction),
303
+ "class": verdict,
304
+ "inference_time": float(inference_time),
305
+ }
306
+
307
+ except Exception as e:
308
+ logger.exception(f"Error during prediction: {e}")
309
+ raise HTTPException(status_code=500, detail=str(e))
310
+
311
+
312
+ if __name__ == "__main__":
313
+ port = int(os.environ.get("MODEL_PORT", 8005))
314
+ 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.0
2
+ torchaudio==2.5.0
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/aasist3/weights/config.json ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "d_args": {
3
+ "architecture": "AASIST",
4
+ "filts": [
5
+ 70,
6
+ [
7
+ 1,
8
+ 32
9
+ ],
10
+ [
11
+ 32,
12
+ 32
13
+ ],
14
+ [
15
+ 32,
16
+ 64
17
+ ],
18
+ [
19
+ 64,
20
+ 64
21
+ ]
22
+ ],
23
+ "first_conv": 128,
24
+ "gat_dims": [
25
+ 64,
26
+ 32
27
+ ],
28
+ "nb_samp": 64600,
29
+ "pool_ratios": [
30
+ 0.5,
31
+ 0.7,
32
+ 0.5,
33
+ 0.5
34
+ ],
35
+ "temperatures": [
36
+ 2.0,
37
+ 2.0,
38
+ 100.0,
39
+ 100.0
40
+ ]
41
+ },
42
+ "load_pretrained": false,
43
+ "size": 200,
44
+ "w2v_cache_dir": "/app/w2v_cache"
45
+ }
audio/aasist3/weights/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a06d43d8d4a7e11b62fb4fa833f13a598fee23be7c368c416df4bec2e44e664f
3
+ size 1287120456
audio/nes2net/.gitignore ADDED
@@ -0,0 +1 @@
 
 
1
+ weights/
audio/nes2net/Dockerfile ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM python:3.10-slim
2
+
3
+ # Install system dependencies
4
+ RUN apt-get update && apt-get install -y \
5
+ git \
6
+ ffmpeg \
7
+ libsndfile1 \
8
+ wget \
9
+ build-essential \
10
+ g++ \
11
+ && rm -rf /var/lib/apt/lists/*
12
+
13
+ WORKDIR /app
14
+
15
+ # Install Python dependencies
16
+ COPY requirements.txt .
17
+ RUN pip install --no-cache-dir -r requirements.txt
18
+
19
+ # Clone fairseq at the specific commit required by Nes2Net,
20
+ # then patch setup.py to skip C extensions (not needed for
21
+ # checkpoint_utils), and install as editable without deps.
22
+ RUN git clone https://github.com/facebookresearch/fairseq.git /app/fairseq_repo && \
23
+ cd /app/fairseq_repo && \
24
+ git checkout a54021305d6b3c4c5959ac9395135f63202db8f1 && \
25
+ sed -i 's/ext_modules=extensions/ext_modules=[]/' setup.py && \
26
+ pip install --no-cache-dir --no-deps -e .
27
+
28
+ # Copy model scripts
29
+ COPY model_scripts /app/model_scripts
30
+
31
+ # Copy API code
32
+ COPY api.py .
33
+
34
+ # Model weights are mounted from the host via docker-compose volumes.
35
+ # Expected weight files in /app/weights/:
36
+ # - nes2net_itw_valaug.pt (fine-tuned Nes2Net checkpoint)
37
+ # - xlsr2_300m.pt (XLSR wav2vec 2.0 for architecture init)
38
+
39
+ EXPOSE 8004
40
+
41
+ # Drop root privileges
42
+ RUN adduser --disabled-password --gecos '' appuser
43
+ USER appuser
44
+
45
+ CMD ["python", "api.py"]
audio/nes2net/api.py ADDED
@@ -0,0 +1,284 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ SAMPLE_RATE = 16000
92
+ TARGET_SAMPLES = 64600 # ~4.04 seconds at 16kHz
93
+
94
+ # Global model instance
95
+ model = None
96
+
97
+
98
+ class AudioInput(BaseModel):
99
+ """Request schema for audio deepfake detection."""
100
+
101
+ audio_data: str = Field(
102
+ ..., description="Base64 encoded audio string (WAV/MP3/etc)"
103
+ )
104
+ threshold: Optional[float] = Field(
105
+ 0.5, ge=0.0, le=1.0, description="Classification threshold"
106
+ )
107
+
108
+
109
+ app = FastAPI(
110
+ title="Nes2Net Audio Deepfake Detection API",
111
+ description=(
112
+ "Service for detecting synthetic speech using the "
113
+ "Nes2Net model (XLSR wav2vec 2.0 + Nested Res2Net TDNN)."
114
+ ),
115
+ version="1.0.0",
116
+ )
117
+
118
+
119
+ def load_model():
120
+ """Load the Nes2Net model with fine-tuned weights.
121
+
122
+ Returns:
123
+ The loaded model, or None if loading fails.
124
+ """
125
+ global model
126
+ if model is not None:
127
+ return model
128
+
129
+ logger.info(f"Loading Nes2Net model onto {DEVICE}...")
130
+
131
+ if Nes2NetModel is None:
132
+ logger.error("Nes2Net model class not available.")
133
+ return None
134
+
135
+ if not os.path.exists(WEIGHTS_PATH):
136
+ logger.error(f"Model weights not found at {WEIGHTS_PATH}")
137
+ return None
138
+
139
+ try:
140
+ args = argparse.Namespace(
141
+ n_output_logits=2,
142
+ dilation=2,
143
+ pool_func="mean",
144
+ SE_ratio=[1],
145
+ Nes_ratio=[8, 8],
146
+ )
147
+ model = Nes2NetModel(args, str(DEVICE))
148
+
149
+ # Load fine-tuned weights
150
+ try:
151
+ state_dict = torch.load(
152
+ WEIGHTS_PATH,
153
+ map_location=DEVICE,
154
+ weights_only=False,
155
+ )
156
+ except TypeError:
157
+ state_dict = torch.load(WEIGHTS_PATH, map_location=DEVICE)
158
+
159
+ model.load_state_dict(state_dict)
160
+ model.to(DEVICE)
161
+ model.eval()
162
+
163
+ logger.info("Nes2Net model loaded successfully.")
164
+ return model
165
+ except Exception as e:
166
+ logger.exception(f"Failed to load Nes2Net model: {e}")
167
+ model = None
168
+ return None
169
+
170
+
171
+ @app.on_event("startup")
172
+ async def startup_event():
173
+ """Load model on service startup."""
174
+ load_model()
175
+
176
+
177
+ @app.get("/health")
178
+ async def health():
179
+ """Health check endpoint."""
180
+ return {
181
+ "status": "healthy" if model is not None else "degraded",
182
+ "model": MODEL_NAME,
183
+ "model_id": MODEL_ID,
184
+ "device": str(DEVICE),
185
+ "weights_found": os.path.exists(WEIGHTS_PATH),
186
+ }
187
+
188
+
189
+ def preprocess_audio(audio_bytes: bytes) -> torch.Tensor:
190
+ """Preprocess audio for Nes2Net inference.
191
+
192
+ Loads audio, resamples to 16kHz mono, and pads/trims
193
+ to TARGET_SAMPLES using tiling (matching original training
194
+ preprocessing).
195
+
196
+ Args:
197
+ audio_bytes: Raw audio file bytes.
198
+
199
+ Returns:
200
+ Audio tensor of shape (1, TARGET_SAMPLES).
201
+
202
+ Raises:
203
+ ValueError: If audio preprocessing fails.
204
+ """
205
+ try:
206
+ logger.info("Starting audio preprocessing...")
207
+ audio, sr = librosa.load(io.BytesIO(audio_bytes), sr=SAMPLE_RATE, mono=True)
208
+ logger.info(f"Audio loaded. Length: {len(audio)} samples at {sr}Hz")
209
+
210
+ # Pad/trim to TARGET_SAMPLES using tiling
211
+ if len(audio) >= TARGET_SAMPLES:
212
+ audio = audio[:TARGET_SAMPLES]
213
+ else:
214
+ num_repeats = TARGET_SAMPLES // len(audio) + 1
215
+ audio = np.tile(audio, num_repeats)[:TARGET_SAMPLES]
216
+
217
+ logger.info(f"Audio padded/trimmed to {TARGET_SAMPLES} samples")
218
+
219
+ audio_tensor = torch.FloatTensor(audio).unsqueeze(0).to(DEVICE)
220
+ return audio_tensor
221
+ except Exception as e:
222
+ logger.error(f"Error preprocessing audio: {e}")
223
+ raise ValueError(f"Audio preprocessing failed: {str(e)}")
224
+
225
+
226
+ @app.post("/predict")
227
+ async def predict(input_data: AudioInput):
228
+ """Run deepfake detection on base64-encoded audio.
229
+
230
+ The model outputs 2 logits: [spoof_score, bonafide_score].
231
+ Class 0 = spoof (fake), Class 1 = bonafide (real).
232
+ The returned probability is the spoof/fake probability.
233
+ """
234
+ if model is None:
235
+ if load_model() is None:
236
+ raise HTTPException(status_code=503, detail="Model not loaded")
237
+
238
+ try:
239
+ start_time = time.time()
240
+ logger.info(
241
+ f"Prediction request. Data size: " f"{len(input_data.audio_data)} chars"
242
+ )
243
+
244
+ # Decode base64 audio
245
+ audio_bytes = base64.b64decode(input_data.audio_data)
246
+
247
+ # Preprocess
248
+ audio_tensor = preprocess_audio(audio_bytes)
249
+
250
+ # Inference
251
+ logger.info("Starting model inference...")
252
+ with torch.no_grad():
253
+ output = model(audio_tensor)
254
+
255
+ # output shape: [batch, 2]
256
+ # Index 0 = spoof logit, Index 1 = bonafide logit
257
+ probs = torch.softmax(output, dim=1)
258
+ prob_fake = probs[0, 0].item()
259
+
260
+ prediction = 1 if prob_fake >= input_data.threshold else 0
261
+ verdict = "fake" if prediction == 1 else "real"
262
+ inference_time = time.time() - start_time
263
+
264
+ logger.info(
265
+ f"Prediction: {verdict} (prob_fake={prob_fake:.4f}, "
266
+ f"time={inference_time:.3f}s)"
267
+ )
268
+
269
+ return {
270
+ "model": MODEL_NAME,
271
+ "probability": float(prob_fake),
272
+ "prediction": int(prediction),
273
+ "class": verdict,
274
+ "inference_time": float(inference_time),
275
+ }
276
+
277
+ except Exception as e:
278
+ logger.exception(f"Error during prediction: {e}")
279
+ raise HTTPException(status_code=500, detail=str(e))
280
+
281
+
282
+ if __name__ == "__main__":
283
+ port = int(os.environ.get("MODEL_PORT", 8004))
284
+ 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.0
2
+ torchaudio==2.5.0
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/nes2net/weights/nes2net_itw_valaug.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:891de3f03cd846a5d6b69a193889440e26a5dde4028c7c4ff3bc80273987a024
3
+ size 1272037256
audio/nes2net/weights/xlsr2_300m.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b08927597f2c9eb2ebd7dcc3ac78ee4b5f6021cbac4b3a6c5a9deec445d80ed9
3
+ size 3808868242
audio/safeear/Dockerfile ADDED
@@ -0,0 +1,39 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM python:3.9-slim
2
+
3
+ RUN apt-get update && apt-get install -y \
4
+ git \
5
+ ffmpeg \
6
+ libsndfile1 \
7
+ wget \
8
+ build-essential \
9
+ && rm -rf /var/lib/apt/lists/*
10
+
11
+ WORKDIR /app
12
+
13
+ COPY requirements.txt .
14
+ RUN pip install --no-cache-dir -r requirements.txt
15
+
16
+ # Clone SafeEar repository (for model code imports)
17
+ RUN git clone --depth 1 https://github.com/LetterLiGo/SafeEar.git /app/safeear_repo
18
+
19
+ # Install the fairseq fork required by SafeEar's SpeechTokenizer
20
+ WORKDIR /app/safeear_repo/fairseq_ours
21
+ RUN pip install --no-cache-dir -e .
22
+ WORKDIR /app
23
+
24
+ # Download model weights from HuggingFace
25
+ RUN mkdir -p /app/weights && \
26
+ wget -q -O /app/weights/SpeechTokenizer.pt \
27
+ "https://huggingface.co/TEC2004/SafeEar-ASV19-spoof-detection/resolve/main/SpeechTokenizer.pt" && \
28
+ wget -q -O /app/weights/model.ckpt \
29
+ "https://huggingface.co/TEC2004/SafeEar-ASV19-spoof-detection/resolve/main/model.ckpt"
30
+
31
+ COPY api.py .
32
+
33
+ EXPOSE 8002
34
+
35
+ # Drop root privileges
36
+ RUN adduser --disabled-password --gecos '' appuser
37
+ USER appuser
38
+
39
+ CMD ["python", "api.py"]
audio/safeear/api.py ADDED
@@ -0,0 +1,290 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ SAMPLE_RATE = 16000
76
+ MAX_AUDIO_LENGTH = 64600 # ~4 seconds at 16kHz (ASVspoof standard)
77
+ SOFTMAX_TEMPERATURE = 5.0 # Calibration temperature for out-of-distribution data
78
+ NUM_INFERENCE_PASSES = 5 # Monte Carlo passes for stable predictions
79
+
80
+ # Global model instances
81
+ decouple_model = None
82
+ detect_model = None
83
+
84
+
85
+ class AudioInput(BaseModel):
86
+ """Schema for audio prediction requests."""
87
+
88
+ audio_data: str = Field(
89
+ ..., description="Base64 encoded audio string (WAV/MP3/etc)"
90
+ )
91
+ threshold: Optional[float] = Field(
92
+ 0.5, ge=0.0, le=1.0, description="Classification threshold"
93
+ )
94
+
95
+
96
+ app = FastAPI(
97
+ title="SafeEar Audio Deepfake Detection API",
98
+ description="Content privacy-preserving deepfake detection using SafeEar.",
99
+ version="1.0.0",
100
+ )
101
+
102
+
103
+ def load_models():
104
+ """Load both the decouple model (SpeechTokenizer) and detect model."""
105
+ global decouple_model, detect_model
106
+
107
+ if decouple_model is not None and detect_model is not None:
108
+ return True
109
+
110
+ speech_tokenizer_path = os.path.join(WEIGHTS_DIR, "SpeechTokenizer.pt")
111
+ checkpoint_path = os.path.join(WEIGHTS_DIR, "model.ckpt")
112
+
113
+ if not os.path.exists(speech_tokenizer_path):
114
+ logger.error(f"SpeechTokenizer weights not found: {speech_tokenizer_path}")
115
+ return False
116
+ if not os.path.exists(checkpoint_path):
117
+ logger.error(f"Model checkpoint not found: {checkpoint_path}")
118
+ return False
119
+
120
+ try:
121
+ # --- Load SpeechTokenizer (decouple model) ---
122
+ from safeear.models.decouple import SpeechTokenizer
123
+
124
+ logger.info("Loading SpeechTokenizer...")
125
+ decouple_model = SpeechTokenizer(
126
+ n_filters=64,
127
+ strides=[8, 5, 4, 2],
128
+ dimension=1024,
129
+ semantic_dimension=768,
130
+ bidirectional=True,
131
+ dilation_base=2,
132
+ residual_kernel_size=3,
133
+ n_residual_layers=1,
134
+ lstm_layers=2,
135
+ activation="ELU",
136
+ codebook_size=1024,
137
+ n_q=8,
138
+ sample_rate=16000,
139
+ )
140
+ st_state = torch.load(speech_tokenizer_path, map_location="cpu")
141
+ decouple_model.load_state_dict(st_state)
142
+ decouple_model.to(DEVICE)
143
+ decouple_model.eval()
144
+ logger.info("SpeechTokenizer loaded.")
145
+
146
+ # --- Load SafeEar1s (detect model) from Lightning checkpoint ---
147
+ from safeear.models.safeear import SafeEar1s, SE_Rawformer_front
148
+
149
+ logger.info("Loading SafeEar1s detect model...")
150
+ detect_model = SafeEar1s(
151
+ front=SE_Rawformer_front(),
152
+ embedding_dim=1024,
153
+ dropout_rate=0.1,
154
+ attention_dropout=0.1,
155
+ stochastic_depth=0.1,
156
+ num_layers=2,
157
+ num_heads=8,
158
+ num_classes=2,
159
+ positional_embedding="sine",
160
+ mlp_ratio=1.0,
161
+ )
162
+
163
+ # The .ckpt is a PyTorch Lightning checkpoint
164
+ ckpt = torch.load(checkpoint_path, map_location="cpu")
165
+ state_dict = ckpt.get("state_dict", ckpt)
166
+
167
+ # Lightning prefixes keys with "detect_model."
168
+ detect_state = {}
169
+ for k, v in state_dict.items():
170
+ if k.startswith("detect_model."):
171
+ detect_state[k.replace("detect_model.", "", 1)] = v
172
+
173
+ detect_model.load_state_dict(detect_state)
174
+ detect_model.to(DEVICE)
175
+ detect_model.eval()
176
+ logger.info("SafeEar1s detect model loaded.")
177
+ return True
178
+
179
+ except Exception as e:
180
+ logger.exception(f"Failed to load SafeEar models: {e}")
181
+ decouple_model = None
182
+ detect_model = None
183
+ return False
184
+
185
+
186
+ def preprocess_audio(audio_bytes: bytes) -> torch.Tensor:
187
+ """Load audio bytes, resample to 16kHz mono, pad/trim."""
188
+ with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
189
+ tmp.write(audio_bytes)
190
+ tmp_path = tmp.name
191
+
192
+ try:
193
+ waveform, _ = librosa.load(tmp_path, sr=SAMPLE_RATE, mono=True)
194
+ finally:
195
+ os.unlink(tmp_path)
196
+
197
+ if len(waveform) < MAX_AUDIO_LENGTH:
198
+ waveform = np.pad(waveform, (0, MAX_AUDIO_LENGTH - len(waveform)))
199
+ else:
200
+ waveform = waveform[:MAX_AUDIO_LENGTH]
201
+
202
+ # Shape: (1, 1, samples) -- batch=1, channels=1, time
203
+ tensor = torch.FloatTensor(waveform).unsqueeze(0).unsqueeze(0).to(DEVICE)
204
+ return tensor
205
+
206
+
207
+ @app.on_event("startup")
208
+ async def startup_event():
209
+ """Attempt to load models at startup."""
210
+ load_models()
211
+
212
+
213
+ @app.get("/health")
214
+ async def health():
215
+ """Return service health status and model availability."""
216
+ models_loaded = decouple_model is not None and detect_model is not None
217
+ return {
218
+ "status": "healthy" if models_loaded else "degraded",
219
+ "model": MODEL_NAME,
220
+ "device": str(DEVICE),
221
+ "weights_found": (
222
+ os.path.exists(os.path.join(WEIGHTS_DIR, "SpeechTokenizer.pt"))
223
+ and os.path.exists(os.path.join(WEIGHTS_DIR, "model.ckpt"))
224
+ ),
225
+ }
226
+
227
+
228
+ @app.post("/predict")
229
+ async def predict(input_data: AudioInput):
230
+ """Run SafeEar inference on base64-encoded audio data."""
231
+ if decouple_model is None or detect_model is None:
232
+ if not load_models():
233
+ raise HTTPException(status_code=503, detail="Models not loaded")
234
+
235
+ try:
236
+ start_time = time.time()
237
+ logger.info(
238
+ "Received prediction request. "
239
+ f"Data size: {len(input_data.audio_data)} chars"
240
+ )
241
+
242
+ audio_bytes = base64.b64decode(input_data.audio_data)
243
+ x_wav = preprocess_audio(audio_bytes)
244
+
245
+ with torch.no_grad():
246
+ # Step 1: Extract acoustic tokens via SpeechTokenizer
247
+ # forward() returns:
248
+ # (reconstructed, commit_loss, semantic_feature, acoustic_tokens)
249
+ # layers=[0,1,2,3,4,5,6,7] means layer 0 goes to
250
+ # semantic_feature; layers 1-7 go to acoustic_tokens list
251
+ _, _, _, acoustic_tokens = decouple_model(
252
+ x_wav, layers=[0, 1, 2, 3, 4, 5, 6, 7]
253
+ )
254
+
255
+ # Step 2: Run detection model with Monte Carlo averaging
256
+ # SafeEar1s uses torch.randperm() in forward, so we average
257
+ # multiple passes for stable predictions
258
+ logit_sum = torch.zeros(1, 2, device=DEVICE)
259
+ for _ in range(NUM_INFERENCE_PASSES):
260
+ raw_logits, _ = detect_model(acoustic_tokens)
261
+ logit_sum += raw_logits
262
+ avg_logits = logit_sum / NUM_INFERENCE_PASSES
263
+
264
+ # Step 3: Get fake probability with temperature-scaled softmax
265
+ # The model produces extreme logits that saturate standard
266
+ # softmax. Temperature scaling preserves discrimination while
267
+ # giving more interpretable probabilities.
268
+ probs = torch.softmax(avg_logits / SOFTMAX_TEMPERATURE, dim=-1)
269
+ prob_fake = probs[0, 1].item()
270
+
271
+ prediction = 1 if prob_fake >= input_data.threshold else 0
272
+ verdict = "fake" if prediction == 1 else "real"
273
+ inference_time = time.time() - start_time
274
+
275
+ return {
276
+ "model": MODEL_NAME,
277
+ "probability": float(prob_fake),
278
+ "prediction": int(prediction),
279
+ "class": verdict,
280
+ "inference_time": float(inference_time),
281
+ }
282
+
283
+ except Exception as e:
284
+ logger.exception(f"Error during prediction: {e}")
285
+ raise HTTPException(status_code=500, detail=str(e))
286
+
287
+
288
+ if __name__ == "__main__":
289
+ port = int(os.environ.get("MODEL_PORT", 8002))
290
+ 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>=1.13.0
2
+ torchaudio>=0.13.0
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/safeear/safeear_repo/.gitignore ADDED
@@ -0,0 +1,167 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Byte-compiled / optimized / DLL files
2
+ __pycache__/
3
+ *.py[cod]
4
+ *$py.class
5
+
6
+ # C extensions
7
+ *.so
8
+
9
+ # Distribution / packaging
10
+ .Python
11
+ build/
12
+ develop-eggs/
13
+ dist/
14
+ downloads/
15
+ eggs/
16
+ .eggs/
17
+ lib/
18
+ lib64/
19
+ parts/
20
+ sdist/
21
+ var/
22
+ wheels/
23
+ share/python-wheels/
24
+ *.egg-info/
25
+ .installed.cfg
26
+ *.egg
27
+ MANIFEST
28
+
29
+ # PyInstaller
30
+ # Usually these files are written by a python script from a template
31
+ # before PyInstaller builds the exe, so as to inject date/other infos into it.
32
+ *.manifest
33
+ *.spec
34
+
35
+ # Installer logs
36
+ pip-log.txt
37
+ pip-delete-this-directory.txt
38
+
39
+ # Unit test / coverage reports
40
+ htmlcov/
41
+ .tox/
42
+ .nox/
43
+ .coverage
44
+ .coverage.*
45
+ .cache
46
+ nosetests.xml
47
+ coverage.xml
48
+ *.cover
49
+ *.py,cover
50
+ .hypothesis/
51
+ .pytest_cache/
52
+ cover/
53
+
54
+ # Translations
55
+ *.mo
56
+ *.pot
57
+
58
+ # Django stuff:
59
+ *.log
60
+ local_settings.py
61
+ db.sqlite3
62
+ db.sqlite3-journal
63
+
64
+ # Flask stuff:
65
+ instance/
66
+ .webassets-cache
67
+
68
+ # Scrapy stuff:
69
+ .scrapy
70
+
71
+ # Sphinx documentation
72
+ docs/_build/
73
+
74
+ # PyBuilder
75
+ .pybuilder/
76
+ target/
77
+
78
+ # Jupyter Notebook
79
+ .ipynb_checkpoints
80
+
81
+ # IPython
82
+ profile_default/
83
+ ipython_config.py
84
+
85
+ # pyenv
86
+ # For a library or package, you might want to ignore these files since the code is
87
+ # intended to run in multiple environments; otherwise, check them in:
88
+ # .python-version
89
+
90
+ # pipenv
91
+ # According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
92
+ # However, in case of collaboration, if having platform-specific dependencies or dependencies
93
+ # having no cross-platform support, pipenv may install dependencies that don't work, or not
94
+ # install all needed dependencies.
95
+ #Pipfile.lock
96
+
97
+ # poetry
98
+ # Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
99
+ # This is especially recommended for binary packages to ensure reproducibility, and is more
100
+ # commonly ignored for libraries.
101
+ # https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
102
+ #poetry.lock
103
+
104
+ # pdm
105
+ # Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
106
+ #pdm.lock
107
+ # pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it
108
+ # in version control.
109
+ # https://pdm.fming.dev/#use-with-ide
110
+ .pdm.toml
111
+
112
+ # PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
113
+ __pypackages__/
114
+
115
+ # Celery stuff
116
+ celerybeat-schedule
117
+ celerybeat.pid
118
+
119
+ # SageMath parsed files
120
+ *.sage.py
121
+
122
+ # Environments
123
+ .env
124
+ .venv
125
+ env/
126
+ venv/
127
+ ENV/
128
+ env.bak/
129
+ venv.bak/
130
+
131
+ # Spyder project settings
132
+ .spyderproject
133
+ .spyproject
134
+
135
+ # Rope project settings
136
+ .ropeproject
137
+
138
+ # mkdocs documentation
139
+ /site
140
+
141
+ # mypy
142
+ .mypy_cache/
143
+ .dmypy.json
144
+ dmypy.json
145
+
146
+ # Pyre type checker
147
+ .pyre/
148
+
149
+ # pytype static type analyzer
150
+ .pytype/
151
+
152
+ # Cython debug symbols
153
+ cython_debug/
154
+
155
+ # PyCharm
156
+ # JetBrains specific template is maintained in a separate JetBrains.gitignore that can
157
+ # be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
158
+ # and can be added to the global gitignore or merged into this file. For a more nuclear
159
+ # option (not recommended) you can uncomment the following to ignore the entire idea folder.
160
+ #.idea/
161
+ model_zoos/*
162
+ Exps/*
163
+ datas/datasets
164
+ datas/ASVSpoof2019/LA
165
+ datas/ASVSpoof2021/ASVspoof2021_LA_eval
166
+ datas/ASVSpoof2021/keys
167
+ create_tsv.py
audio/safeear/safeear_repo/LICENSE ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Creative Commons Attribution 4.0 International License
2
+
3
+ ## License
4
+
5
+ You are free to:
6
+
7
+ - Share — copy and redistribute the material in any medium or format
8
+ - Adapt — remix, transform, and build upon the material for any purpose, even commercially.
9
+
10
+ Under the following terms:
11
+
12
+ 1. **Attribution** — You must give appropriate credit, provide a link to the license, and indicate if changes were made. You may do so in any reasonable manner, but not in any way that suggests the licensor endorses you or your use.
13
+
14
+ 2. **No additional restrictions** — You may not apply legal terms or technological measures that legally restrict others from doing anything the license permits.
15
+
16
+ ## Other Terms
17
+
18
+ - This license applies to all types of works, including but not limited to text, images, audio, video, etc.
19
+ - This license does not apply to any third-party materials included in the work, for which you must obtain permission separately.
20
+
21
+ ## Disclaimer
22
+
23
+ This work is provided on an "as is" basis, without any warranties or conditions of any kind, either express or implied, including but not limited to implied warranties of merchantability, fitness for a particular purpose, or non-infringement.
audio/safeear/safeear_repo/README.md ADDED
@@ -0,0 +1,134 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # <font color=E7595C>Safe</font><font color=F6C446>Ear</font><img src="assert/SafeEar_logo.jpg" alt="icon" style="width: 2em; height: 1.5em; vertical-align: middle;">: <font color=E7595C>Content Privacy-Preserving</font> <font color=F6C446>Audio Deepfake Detection</font>
2
+
3
+ [![arXiv](https://img.shields.io/badge/arXiv-2409.09272-b31b1b.svg)](https://arxiv.org/abs/2409.09272)
4
+ [![PRs Welcome](https://img.shields.io/badge/PRs-welcome-brightgreen.svg?style=flat-square)](https://makeapullrequest.com)
5
+ [![CC BY 4.0](https://img.shields.io/badge/license-CC%20BY%204.0-blue.svg)](https://creativecommons.org/licenses/by/4.0/)
6
+ ![GitHub stars](https://img.shields.io/github/stars/LetterLiGo/SafeEar)
7
+ ![GitHub forks](https://img.shields.io/github/forks/LetterLiGo/SafeEar)
8
+ ![Website](https://img.shields.io/website?url=https://safeearweb.github.io/Project/)
9
+
10
+
11
+ By [1] Zhejiang University, [2] Tsinghua University.
12
+ * [Xinfeng Li](https://letterligo.github.io)* [1], [Kai Li](https://cslikai.cn)* [2], Yifan Zheng [1], Chen Yan† [1], Xiaoyu Ji [1], Wenyuan Xu [1].
13
+
14
+ This repository is an official implementation of the SafeEar accepted to **ACM CCS 2024** (Core-A*, CCF-A, Big4) .
15
+
16
+ Please also visit our <a href="https://safeearweb.github.io/Project/">(1) Project Website</a>, <a href="https://zenodo.org/records/14062964">(2) Full CVoiceFake Dataset</a>, and <a href="https://zenodo.org/records/11124319">(3) Sampled CVoiceFake Dataset</a>.
17
+
18
+ ## 🔥News
19
+
20
+ [2025-03-18]: Supported the batch testing for ASVspoof 2019 and 2021, fixed some bugs for datasets and trainer.
21
+
22
+ [2024-12-10]: Fixed all the bugs for training and test, and uploaded the files for data generation `datas/`.
23
+
24
+ [2024-12-01]: Uploaded the checkpoint for data generation `datas/`.
25
+
26
+ ## ✨Key Highlights:
27
+
28
+ In this paper, we propose SafeEar, a novel framework that aims to detect deepfake audios without relying on accessing the speech content within. Our key idea is to devise a neural audio codec into a novel decoupling model that well separates the semantic and acoustic information from audio samples, and only use the acoustic information (e.g., prosody and timbre) for deepfake detection. In this way, no semantic content will be exposed to the detector. To overcome the challenge of identifying diverse deepfake audio without semantic clues, we enhance our deepfake detector with multi-head self-attention and codec augmentation. Extensive experiments conducted on four benchmark datasets demonstrate SafeEar’s effectiveness in detecting various deepfake techniques with an equal error rate (EER) down to 2.02%. Simultaneously, it shields five-language speech content from being deciphered by both machine and human auditory analysis, demonstrated by word error rates (WERs) all above 93.93% and our user study. Furthermore, our benchmark constructed for anti-deepfake and anti-content recovery evaluation helps provide a basis for future research in the realms of audio privacy preservation and deepfake detection.
29
+
30
+ ## 🚀Overall Pipeline
31
+
32
+ ![pipeline](assert/overall.gif)
33
+
34
+ ## 🔧Installation
35
+
36
+ 1. Clone the repository:
37
+
38
+ ```shell
39
+ git clone git@github.com:LetterLiGo/SafeEar.git
40
+ cd SafeEar/
41
+ ```
42
+
43
+ 2. Create and activate the conda environment:
44
+
45
+ ```shell
46
+ conda create -n safeear python=3.9
47
+ conda activate safeear
48
+ ```
49
+
50
+ 3. Install PyTorch and torchvision following the [official instructions](https://pytorch.org). The code requires `python=3.9`, `pytorch=1.13`, `torchvision=0.14`.
51
+
52
+
53
+ ```shell
54
+ pip install torch==1.13.1+cu116 torchvision==0.14.1+cu116 torchaudio==0.13.1 --extra-index-url https://download.pytorch.org/whl/cu116
55
+
56
+ ```
57
+ 4. Install other dependencies:
58
+
59
+ ```shell
60
+ pip install pip==24.0
61
+ pip install -r requirements.txt
62
+ ```
63
+
64
+ ## 📊Model Performance
65
+ ### ASVspoof 2019 & 2021
66
+ ![](assert/ASVSpoof-results.png)
67
+ ### Speech Recognition Performance
68
+ ![](assert/exp1.png)
69
+
70
+ ## Data preparation
71
+
72
+ ### AVSpoof 2019 & 2021
73
+
74
+ Please download the [ASVspoof 2019](https://datashare.is.ed.ac.uk/handle/10283/3336) and [ASVspoof 2021](https://www.asvspoof.org/index2021.html) datasets and extract them to the `datas/datasets` directory.
75
+
76
+ ```shell
77
+ datas/datasets/ASVspoof2019
78
+ datas/datasets/ASVspoof2021
79
+ ```
80
+
81
+ #### Generate the Hubert L9 feature files
82
+
83
+ ```shell
84
+ mkdir model_zoos
85
+ cd model_zoos
86
+ wget https://dl.fbaipublicfiles.com/hubert/hubert_base_ls960.pt
87
+ wget https://cloud.tsinghua.edu.cn/f/413a0cd2e6f749eea956/?dl=1 -O SpeechTokenizer.pt
88
+ cd ../datas
89
+ # Generate the Hubert L9 feature files for ASVspoof 2019
90
+ python dump_hubert_avg_feature.py datasets/ASVSpoof2019 datasets/ASVSpoof2019_Hubert_L9
91
+ # Generate the Hubert L9 feature files for ASVspoof 2021
92
+ python dump_hubert_avg_feature.py datasets/ASVSpoof2021 datasets/ASVSpoof2021_Hubert_L9
93
+ ```
94
+
95
+ ## 📚Training
96
+
97
+ Before starting training, please modify the parameter configurations in [`configs`](configs).
98
+
99
+ Use the following commands to start training:
100
+
101
+ ```shell
102
+ python train.py --conf_dir config/train19.yaml
103
+ python train.py --conf_dir config/train21.yaml
104
+ ```
105
+
106
+ ## 📈Testing/Inference
107
+
108
+ To evaluate a model on one or more GPUs, specify the `CUDA_VISIBLE_DEVICES`, `dataset`, `model` and `checkpoint`:
109
+
110
+ ```shell
111
+ python test.py --conf_dir Exps/ASVspoof19/config.yaml
112
+ python test.py --conf_dir Exps/ASVspoof21/config.yaml
113
+ ```
114
+
115
+ ## Bugs and Issues
116
+
117
+ If you meet `RuntimeError: Failed to load audio from <_io.BytesIO object at 0x7f45cb978f90>`, please use the following command to fix it:
118
+
119
+ ```shell
120
+ conda install -c anaconda 'ffmpeg<4.4'
121
+ ```
122
+
123
+ ## 📜Citation
124
+
125
+ If you find our work/code/dataset helpful, please consider citing:
126
+
127
+ ```
128
+ @inproceedings{li2024safeear,
129
+ author = {Li, Xinfeng and Li, Kai and Zheng, Yifan and Yan, Chen and Ji, Xiaoyu, and Xu, Wenyuan},
130
+ title = {{SafeEar: Content Privacy-Preserving Audio Deepfake Detection}},
131
+ booktitle = {Proceedings of the 2024 {ACM} {SIGSAC} Conference on Computer and Communications Security (CCS)}
132
+ year = {2024},
133
+ }
134
+ ```
audio/safeear/safeear_repo/assert/ASVSpoof-results.png ADDED

Git LFS Details

  • SHA256: 0697de878b12c7b96ad4609569ddc20388027764ca441575b9c66dc6696d53d2
  • Pointer size: 131 Bytes
  • Size of remote file: 684 kB
audio/safeear/safeear_repo/assert/Fig1.jpg ADDED

Git LFS Details

  • SHA256: 433d57f7a1496c0b3ad6c6c06434261ba2b98901a2fccbd821326a8637fc8135
  • Pointer size: 131 Bytes
  • Size of remote file: 236 kB
audio/safeear/safeear_repo/assert/SafeEar_logo.jpg ADDED

Git LFS Details

  • SHA256: e9ab20f08e487b0454686259bf9c52a4b79d35144e56a40f3f8abd94ce0e637f
  • Pointer size: 131 Bytes
  • Size of remote file: 258 kB
audio/safeear/safeear_repo/assert/exp1.png ADDED

Git LFS Details

  • SHA256: 5f6468a63356b243039324d80e928936d5eaac34b668e4410090b8418f5d0918
  • Pointer size: 132 Bytes
  • Size of remote file: 1.64 MB
audio/safeear/safeear_repo/assert/overall.gif ADDED

Git LFS Details

  • SHA256: a946cf1f760f5f94cdf8907ff2ac70f38ea211964c9df6c1e9a2f1237be6aed0
  • Pointer size: 131 Bytes
  • Size of remote file: 114 kB
audio/safeear/safeear_repo/assert/overall.mp4 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d34cf40c912995aeb9070d168bc41663e365b5e64544538e8cb7e36770233afc
3
+ size 2170657
audio/safeear/safeear_repo/assert/overall.png ADDED

Git LFS Details

  • SHA256: f0271961e7957d83622d2b24a0f6f791bd131a1a493a72abc1142fd8b4a48e1a
  • Pointer size: 131 Bytes
  • Size of remote file: 121 kB
audio/safeear/safeear_repo/assert/safe-space.png ADDED
audio/safeear/safeear_repo/config/train19.yaml ADDED
@@ -0,0 +1,87 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ datamodule:
2
+ _target_: safeear.datas.asvspoof19.DataModule
3
+ batch_size: 2
4
+ num_workers: 8
5
+ pin_memory: true
6
+ DataClass_dict:
7
+ _target_: safeear.datas.asvspoof19.DataClass
8
+ train_path: ["datas/ASVSpoof2019/train.tsv", "datas/ASVSpoof2019/ASVspoof2019.LA.cm.train.trn.txt", "datas/datasets/ASVSpoof2019_Hubert_L9/ASVspoof2019_LA_train/flac"]
9
+ val_path: ["datas/ASVSpoof2019/dev.tsv", "datas/ASVSpoof2019/ASVspoof2019.LA.cm.dev.trl.txt", "datas/datasets/ASVSpoof2019_Hubert_L9/ASVspoof2019_LA_dev/flac"]
10
+ test_path: ["datas/ASVSpoof2019/eval.tsv", "datas/ASVSpoof2019/ASVspoof2019.LA.cm.eval.trl.txt", "datas/datasets/ASVSpoof2019_Hubert_L9/ASVspoof2019_LA_eval/flac"]
11
+ max_len: 64600
12
+
13
+ decouple_model:
14
+ _target_: safeear.models.decouple.SpeechTokenizer
15
+ n_filters: 64
16
+ strides: [8,5,4,2]
17
+ dimension: 1024
18
+ semantic_dimension: 768
19
+ bidirectional: true
20
+ dilation_base: 2
21
+ residual_kernel_size: 3
22
+ n_residual_layers: 1
23
+ lstm_layers: 2
24
+ activation: ELU
25
+ codebook_size: 1024
26
+ n_q: 8
27
+ sample_rate: 16000
28
+
29
+ speechtokenizer_path: model_zoos/SpeechTokenizer.pt
30
+
31
+ detect_model:
32
+ _target_: safeear.models.safeear.SafeEar1s
33
+ front:
34
+ _target_: safeear.models.safeear.SE_Rawformer_front
35
+ embedding_dim: 1024
36
+ dropout_rate: 0.1
37
+ attention_dropout: 0.1
38
+ stochastic_depth: 0.1
39
+ num_layers: 2
40
+ num_heads: 8
41
+ num_classes: 2
42
+ positional_embedding: 'sine'
43
+ mlp_ratio: 1.0
44
+
45
+ system:
46
+ _target_: safeear.trainer.safeear_trainer.SafeEarTrainer
47
+ lr_raw_former: 3.0e-4
48
+ save_score_path: ${exp.dir}/${exp.name}
49
+
50
+ exp:
51
+ dir: Exps/ # 修改
52
+ name: ASVspoof19 # 修改
53
+
54
+ early_stopping:
55
+ _target_: pytorch_lightning.callbacks.EarlyStopping
56
+ monitor: val_eer # 修改
57
+ mode: min
58
+ patience: 40
59
+ verbose: true
60
+
61
+ checkpoint:
62
+ _target_: pytorch_lightning.callbacks.ModelCheckpoint
63
+ dirpath: ${exp.dir}/${exp.name}/checkpoints
64
+ monitor: val_eer # 修改
65
+ mode: min
66
+ verbose: true
67
+ save_top_k: 1
68
+ save_last: true
69
+ filename: '{epoch}-{val_eer:.4f}' # 修改
70
+
71
+ logger:
72
+ _target_: pytorch_lightning.loggers.WandbLogger
73
+ name: ${exp.name}
74
+ save_dir: ${exp.dir}/${exp.name}/logs
75
+ offline: true
76
+ project: SafeEar
77
+
78
+ trainer:
79
+ _target_: pytorch_lightning.Trainer
80
+ devices: [0]
81
+ max_epochs: 500
82
+ sync_batchnorm: true
83
+ default_root_dir: ${exp.dir}/${exp.name}/
84
+ accelerator: gpu
85
+ limit_train_batches: 1.0
86
+ limit_val_batches: 1.0
87
+ fast_dev_run: false
audio/safeear/safeear_repo/config/train21.yaml ADDED
@@ -0,0 +1,87 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ datamodule:
2
+ _target_: safeear.datas.asvspoof21.DataModule
3
+ batch_size: 2
4
+ num_workers: 8
5
+ pin_memory: true
6
+ DataClass_dict:
7
+ _target_: safeear.datas.asvspoof21.DataClass
8
+ train_path: ["datas/ASVSpoof2019/train.tsv", "datas/ASVSpoof2019/ASVspoof2019.LA.cm.train.trn.txt", "datas/datasets/ASVSpoof2019_Hubert_L9/ASVspoof2019_LA_train/flac"]
9
+ val_path: ["datas/ASVSpoof2019/dev.tsv", "datas/ASVSpoof2019/ASVspoof2019.LA.cm.dev.trl.txt", "datas/datasets/ASVSpoof2019_Hubert_L9/ASVspoof2019_LA_dev/flac"]
10
+ test_path: ["datas/ASVSpoof2021/eval.tsv", "datas/ASVSpoof2021/ASVspoof2021.LA.cm.eval.trl.txt", "datas/datasets/ASVSpoof2021_Hubert_L9"]
11
+ max_len: 64600
12
+
13
+ decouple_model:
14
+ _target_: safeear.models.decouple.SpeechTokenizer
15
+ n_filters: 64
16
+ strides: [8,5,4,2]
17
+ dimension: 1024
18
+ semantic_dimension: 768
19
+ bidirectional: true
20
+ dilation_base: 2
21
+ residual_kernel_size: 3
22
+ n_residual_layers: 1
23
+ lstm_layers: 2
24
+ activation: ELU
25
+ codebook_size: 1024
26
+ n_q: 8
27
+ sample_rate: 16000
28
+
29
+ speechtokenizer_path: model_zoos/SpeechTokenizer.pt
30
+
31
+ detect_model:
32
+ _target_: safeear.models.safeear.SafeEar1s
33
+ front:
34
+ _target_: safeear.models.safeear.SE_Rawformer_front
35
+ embedding_dim: 1024
36
+ dropout_rate: 0.1
37
+ attention_dropout: 0.1
38
+ stochastic_depth: 0.1
39
+ num_layers: 2
40
+ num_heads: 8
41
+ num_classes: 2
42
+ positional_embedding: 'sine'
43
+ mlp_ratio: 1.0
44
+
45
+ system:
46
+ _target_: safeear.trainer.safeear_trainer.SafeEarTrainer
47
+ lr_raw_former: 3.0e-4
48
+ save_score_path: ${exp.dir}/${exp.name}
49
+
50
+ exp:
51
+ dir: Exps/ # 修改
52
+ name: ASVspoof21 # 修改
53
+
54
+ early_stopping:
55
+ _target_: pytorch_lightning.callbacks.EarlyStopping
56
+ monitor: val_eer # 修改
57
+ mode: min
58
+ patience: 40
59
+ verbose: true
60
+
61
+ checkpoint:
62
+ _target_: pytorch_lightning.callbacks.ModelCheckpoint
63
+ dirpath: ${exp.dir}/${exp.name}/checkpoints
64
+ monitor: val_eer # 修改
65
+ mode: min
66
+ verbose: true
67
+ save_top_k: 1
68
+ save_last: true
69
+ filename: '{epoch}-{val_eer:.4f}' # 修改
70
+
71
+ logger:
72
+ _target_: pytorch_lightning.loggers.WandbLogger
73
+ name: ${exp.name}
74
+ save_dir: ${exp.dir}/${exp.name}/logs
75
+ offline: true
76
+ project: SafeEar
77
+
78
+ trainer:
79
+ _target_: pytorch_lightning.Trainer
80
+ devices: [0]
81
+ max_epochs: 40
82
+ sync_batchnorm: true
83
+ default_root_dir: ${exp.dir}/${exp.name}/
84
+ accelerator: gpu
85
+ limit_train_batches: 1.0
86
+ limit_val_batches: 1.0
87
+ fast_dev_run: false
audio/safeear/safeear_repo/datas/ASVSpoof2019/ASVspoof2019.LA.cm.dev.trl.txt ADDED
The diff for this file is too large to render. See raw diff
 
audio/safeear/safeear_repo/datas/ASVSpoof2019/ASVspoof2019.LA.cm.eval.trl.txt ADDED
The diff for this file is too large to render. See raw diff
 
audio/safeear/safeear_repo/datas/ASVSpoof2019/ASVspoof2019.LA.cm.train.trn.txt ADDED
The diff for this file is too large to render. See raw diff
 
audio/safeear/safeear_repo/datas/ASVSpoof2019/dev.tsv ADDED
The diff for this file is too large to render. See raw diff
 
audio/safeear/safeear_repo/datas/ASVSpoof2019/eval.tsv ADDED
The diff for this file is too large to render. See raw diff
 
audio/safeear/safeear_repo/datas/ASVSpoof2019/train.tsv ADDED
The diff for this file is too large to render. See raw diff
 
audio/safeear/safeear_repo/datas/ASVSpoof2021/ASVspoof2021.LA.cm.eval.trl.txt ADDED
The diff for this file is too large to render. See raw diff
 
audio/safeear/safeear_repo/datas/ASVSpoof2021/eval.tsv ADDED
The diff for this file is too large to render. See raw diff
 
audio/safeear/safeear_repo/datas/dump_hubert_avg_feature.py ADDED
@@ -0,0 +1,109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Facebook, Inc. and its affiliates.
2
+ #
3
+ # This source code is licensed under the MIT license found in the
4
+ # LICENSE file in the root directory of this source tree.
5
+
6
+ import logging
7
+ import os
8
+ import sys
9
+ import warnings
10
+
11
+ warnings.filterwarnings('ignore')
12
+ from pathlib import Path
13
+
14
+ import fairseq
15
+ import librosa
16
+ import soundfile as sf
17
+ import torch
18
+ import torch.nn.functional as F
19
+ import tqdm
20
+ from feature_utils import dump_feature, get_path_iterator
21
+ from npy_append_array import NpyAppendArray
22
+
23
+ logging.basicConfig(
24
+ format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
25
+ datefmt="%Y-%m-%d %H:%M:%S",
26
+ level=os.environ.get("LOGLEVEL", "INFO").upper(),
27
+ stream=sys.stdout,
28
+ )
29
+ logger = logging.getLogger("dump_hubert_feature")
30
+
31
+
32
+ class HubertFeatureReader(object):
33
+ def __init__(self, ckpt_path, layer, max_chunk=1600000):
34
+ (
35
+ model,
36
+ cfg,
37
+ task,
38
+ ) = fairseq.checkpoint_utils.load_model_ensemble_and_task([ckpt_path])
39
+ self.model = model[0].eval().cuda()
40
+ self.task = task
41
+ self.layer = layer
42
+ self.max_chunk = max_chunk
43
+ logger.info(f"TASK CONFIG:\n{self.task.cfg}")
44
+ logger.info(f" max_chunk = {self.max_chunk}")
45
+
46
+ def read_audio(self, path, ref_len=None):
47
+ wav, sr = librosa.load(path, sr=None)
48
+ assert sr == self.task.cfg.sample_rate, sr
49
+ if wav.ndim == 2:
50
+ wav = wav.mean(-1)
51
+ assert wav.ndim == 1, wav.ndim
52
+ if ref_len is not None and abs(ref_len - len(wav)) > 160:
53
+ logging.warning(f"ref {ref_len} != read {len(wav)} ({path})")
54
+ return wav
55
+
56
+ def get_feats(self, path, ref_len=None):
57
+ x = self.read_audio(path, ref_len)
58
+ with torch.no_grad():
59
+ x = torch.from_numpy(x).float().cuda()
60
+ if self.task.cfg.normalize:
61
+ x = F.layer_norm(x, x.shape)
62
+ x = x.view(1, -1)
63
+
64
+ avg_feat = []
65
+ for start in range(0, x.size(1), self.max_chunk):
66
+ x_chunk = x[:, start: start + self.max_chunk]
67
+ feat_chunk, _, avg_feat_chunk = self.model.extract_features(
68
+ source=x_chunk,
69
+ padding_mask=None,
70
+ mask=False,
71
+ output_layer=self.layer,
72
+ )
73
+ avg_feat.append(avg_feat_chunk)
74
+ return torch.cat(avg_feat, 1).squeeze(0)
75
+
76
+ def dump_feature(reader,audio_dir,save_dir):
77
+ save_dir = Path(save_dir)
78
+ audio_dir = Path(audio_dir)
79
+
80
+ audio_files = list(audio_dir.glob("**/*.flac"))
81
+ for audio_file in tqdm.tqdm(audio_files):
82
+ releative_path = audio_file.relative_to(audio_dir).with_suffix(".npy")
83
+ save_path = save_dir / releative_path
84
+ if not save_path.parent.exists():
85
+ save_path.parent.mkdir(parents=True)
86
+
87
+ feat_f = NpyAppendArray(save_path)
88
+ feat = reader.get_feats(audio_file)
89
+ feat_f.append(feat.cpu().numpy())
90
+ logger.info("finished successfully")
91
+
92
+ def main(audio_dir, save_dir, ckpt_path, layer, max_chunk):
93
+ reader = HubertFeatureReader(ckpt_path, layer, max_chunk)
94
+ dump_feature(reader, audio_dir, save_dir)
95
+
96
+
97
+ if __name__ == "__main__":
98
+ import argparse
99
+
100
+ parser = argparse.ArgumentParser()
101
+ parser.add_argument("audio_dir", nargs="?", default="datasets/ASVSpoof2019", help="Directory containing audio files")
102
+ parser.add_argument("save_dir", nargs="?", default="datasets/ASVSpoof2019_Hubert_L9", help="Directory to save extracted features")
103
+ parser.add_argument("ckpt_path", nargs="?", default="../model_zoos/hubert_base_ls960.pt", help="Path to the checkpoint file")
104
+ parser.add_argument("layer", nargs="?", type=int, default=9, help="Layer number to extract features from")
105
+ parser.add_argument("--max_chunk", type=int, default=1600000, help="Maximum chunk size for processing")
106
+ args = parser.parse_args()
107
+ logger.info(args)
108
+
109
+ main(**vars(args))