Instructions to use deepsafe/deepsafe-services with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use deepsafe/deepsafe-services with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("deepsafe/deepsafe-services", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Upload DeepSafe services: model weights and inference code
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +94 -0
- audio/aasist3/.gitignore +1 -0
- audio/aasist3/Dockerfile +33 -0
- audio/aasist3/api.py +314 -0
- audio/aasist3/model/__init__.py +1 -0
- audio/aasist3/model/branch.py +34 -0
- audio/aasist3/model/full_model.py +139 -0
- audio/aasist3/model/gat.py +99 -0
- audio/aasist3/model/hs_gal.py +176 -0
- audio/aasist3/model/kan.py +213 -0
- audio/aasist3/model/pool.py +45 -0
- audio/aasist3/model/residual.py +56 -0
- audio/aasist3/model/wav2vec.py +82 -0
- audio/aasist3/requirements.txt +11 -0
- audio/aasist3/weights/config.json +45 -0
- audio/aasist3/weights/model.safetensors +3 -0
- audio/nes2net/.gitignore +1 -0
- audio/nes2net/Dockerfile +45 -0
- audio/nes2net/api.py +284 -0
- audio/nes2net/model_scripts/__init__.py +0 -0
- audio/nes2net/model_scripts/wav2vec2_Nes2Net_X.py +317 -0
- audio/nes2net/requirements.txt +13 -0
- audio/nes2net/weights/nes2net_itw_valaug.pt +3 -0
- audio/nes2net/weights/xlsr2_300m.pt +3 -0
- audio/safeear/Dockerfile +39 -0
- audio/safeear/api.py +290 -0
- audio/safeear/download_weights.sh +23 -0
- audio/safeear/requirements.txt +15 -0
- audio/safeear/safeear_repo/.gitignore +167 -0
- audio/safeear/safeear_repo/LICENSE +23 -0
- audio/safeear/safeear_repo/README.md +134 -0
- audio/safeear/safeear_repo/assert/ASVSpoof-results.png +3 -0
- audio/safeear/safeear_repo/assert/Fig1.jpg +3 -0
- audio/safeear/safeear_repo/assert/SafeEar_logo.jpg +3 -0
- audio/safeear/safeear_repo/assert/exp1.png +3 -0
- audio/safeear/safeear_repo/assert/overall.gif +3 -0
- audio/safeear/safeear_repo/assert/overall.mp4 +3 -0
- audio/safeear/safeear_repo/assert/overall.png +3 -0
- audio/safeear/safeear_repo/assert/safe-space.png +0 -0
- audio/safeear/safeear_repo/config/train19.yaml +87 -0
- audio/safeear/safeear_repo/config/train21.yaml +87 -0
- audio/safeear/safeear_repo/datas/ASVSpoof2019/ASVspoof2019.LA.cm.dev.trl.txt +0 -0
- audio/safeear/safeear_repo/datas/ASVSpoof2019/ASVspoof2019.LA.cm.eval.trl.txt +0 -0
- audio/safeear/safeear_repo/datas/ASVSpoof2019/ASVspoof2019.LA.cm.train.trn.txt +0 -0
- audio/safeear/safeear_repo/datas/ASVSpoof2019/dev.tsv +0 -0
- audio/safeear/safeear_repo/datas/ASVSpoof2019/eval.tsv +0 -0
- audio/safeear/safeear_repo/datas/ASVSpoof2019/train.tsv +0 -0
- audio/safeear/safeear_repo/datas/ASVSpoof2021/ASVspoof2021.LA.cm.eval.trl.txt +0 -0
- audio/safeear/safeear_repo/datas/ASVSpoof2021/eval.tsv +0 -0
- 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 |
+
[](https://arxiv.org/abs/2409.09272)
|
| 4 |
+
[](https://makeapullrequest.com)
|
| 5 |
+
[](https://creativecommons.org/licenses/by/4.0/)
|
| 6 |
+

|
| 7 |
+

|
| 8 |
+

|
| 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 |
+

|
| 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 |
+

|
| 67 |
+
### Speech Recognition Performance
|
| 68 |
+

|
| 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
|
audio/safeear/safeear_repo/assert/Fig1.jpg
ADDED
|
Git LFS Details
|
audio/safeear/safeear_repo/assert/SafeEar_logo.jpg
ADDED
|
Git LFS Details
|
audio/safeear/safeear_repo/assert/exp1.png
ADDED
|
Git LFS Details
|
audio/safeear/safeear_repo/assert/overall.gif
ADDED
|
Git LFS Details
|
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
|
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))
|