deepsafe commited on
Commit
9e14838
·
verified ·
1 Parent(s): 1583ce9

Add stripped inference-only model code mirror

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. clean/audio/nes2net/SOURCE.md +17 -0
  2. clean/audio/nes2net/__init__.py +0 -0
  3. clean/audio/nes2net/wav2vec2_Nes2Net_X.py +317 -0
  4. clean/audio/safeear/.gitignore +167 -0
  5. clean/audio/safeear/LICENSE +23 -0
  6. clean/audio/safeear/README.md +134 -0
  7. clean/audio/safeear/SOURCE.md +17 -0
  8. clean/audio/safeear/config/train19.yaml +87 -0
  9. clean/audio/safeear/config/train21.yaml +87 -0
  10. clean/audio/safeear/requirements.txt +123 -0
  11. clean/audio/safeear/safeear/losses/loss.py +215 -0
  12. clean/audio/safeear/safeear/models/decouple.py +207 -0
  13. clean/audio/safeear/safeear/models/discriminator.py +422 -0
  14. clean/audio/safeear/safeear/models/modules/__init__.py +21 -0
  15. clean/audio/safeear/safeear/models/modules/conv.py +252 -0
  16. clean/audio/safeear/safeear/models/modules/lstm.py +32 -0
  17. clean/audio/safeear/safeear/models/modules/norm.py +28 -0
  18. clean/audio/safeear/safeear/models/modules/quantization/__init__.py +8 -0
  19. clean/audio/safeear/safeear/models/modules/quantization/ac.py +292 -0
  20. clean/audio/safeear/safeear/models/modules/quantization/core_vq.py +366 -0
  21. clean/audio/safeear/safeear/models/modules/quantization/distrib.py +126 -0
  22. clean/audio/safeear/safeear/models/modules/quantization/vq.py +108 -0
  23. clean/audio/safeear/safeear/models/modules/seanet.py +275 -0
  24. clean/audio/safeear/safeear/models/safeear.py +959 -0
  25. clean/audio/safeear/safeear/trainer/safeear_trainer.py +189 -0
  26. clean/audio/safeear/safeear/utils/dump_hubert_feature.py +108 -0
  27. clean/audio/safeear/test.py +78 -0
  28. clean/audio/safeear/train.py +102 -0
  29. clean/audio/shiftyspeech/.env +2 -0
  30. clean/audio/shiftyspeech/LICENSE +21 -0
  31. clean/audio/shiftyspeech/RawBoost.py +143 -0
  32. clean/audio/shiftyspeech/SOURCE.md +17 -0
  33. clean/audio/shiftyspeech/Simplified_CM_solution.py +227 -0
  34. clean/audio/shiftyspeech/data_utils.py +292 -0
  35. clean/audio/shiftyspeech/model.py +603 -0
  36. clean/audio/shiftyspeech/startup_config.py +60 -0
  37. clean/audio/shiftyspeech/train.py +446 -0
  38. clean/image/aide/LICENSE +21 -0
  39. clean/image/aide/README.md +168 -0
  40. clean/image/aide/SOURCE.md +17 -0
  41. clean/image/aide/data/__init__.py +0 -0
  42. clean/image/aide/data/dct.py +107 -0
  43. clean/image/aide/engine_finetune.py +197 -0
  44. clean/image/aide/main_finetune.py +449 -0
  45. clean/image/aide/models/AIDE.py +298 -0
  46. clean/image/aide/models/__init__.py +0 -0
  47. clean/image/aide/models/srm_filter_kernel.py +220 -0
  48. clean/image/aide/models/utils.py +116 -0
  49. clean/image/aide/optim_factory.py +222 -0
  50. clean/image/aide/requirements.txt +37 -0
clean/audio/nes2net/SOURCE.md ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Source: audio/nes2net
2
+
3
+ | Field | Value |
4
+ |---|---|
5
+ | Upstream | **UNVERIFIED** -- provenance was lost when this code was vendored |
6
+ | Paper | not recorded |
7
+ | Commit SHA | **not recorded** -- the vendoring step did not preserve it |
8
+ | Mirrored on | 2026-09-15 |
9
+ | Upstream license | no license file present upstream (all rights reserved) |
10
+
11
+ This is a **mirror**, stripped to the files needed for inference. The full
12
+ untouched snapshot is at `archive/audio__nes2net.tar.gz`.
13
+
14
+ This code is the work of its original authors and is **not** covered by the
15
+ DeepSafe project license. If you are an author and want this removed, open an
16
+ issue on https://github.com/deepsafehq/deepsafe-bench and it will be taken down
17
+ within 48 hours, no questions asked.
clean/audio/nes2net/__init__.py ADDED
File without changes
clean/audio/nes2net/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)
clean/audio/safeear/.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
clean/audio/safeear/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.
clean/audio/safeear/README.md ADDED
@@ -0,0 +1,134 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # <font color=E7595C>Safe</font><font color=F6C446>Ear</font><img src="assert/SafeEar_logo.jpg" alt="icon" style="width: 2em; height: 1.5em; vertical-align: middle;">: <font color=E7595C>Content Privacy-Preserving</font> <font color=F6C446>Audio Deepfake Detection</font>
2
+
3
+ [![arXiv](https://img.shields.io/badge/arXiv-2409.09272-b31b1b.svg)](https://arxiv.org/abs/2409.09272)
4
+ [![PRs Welcome](https://img.shields.io/badge/PRs-welcome-brightgreen.svg?style=flat-square)](https://makeapullrequest.com)
5
+ [![CC BY 4.0](https://img.shields.io/badge/license-CC%20BY%204.0-blue.svg)](https://creativecommons.org/licenses/by/4.0/)
6
+ ![GitHub stars](https://img.shields.io/github/stars/LetterLiGo/SafeEar)
7
+ ![GitHub forks](https://img.shields.io/github/forks/LetterLiGo/SafeEar)
8
+ ![Website](https://img.shields.io/website?url=https://safeearweb.github.io/Project/)
9
+
10
+
11
+ By [1] Zhejiang University, [2] Tsinghua University.
12
+ * [Xinfeng Li](https://letterligo.github.io)* [1], [Kai Li](https://cslikai.cn)* [2], Yifan Zheng [1], Chen Yan† [1], Xiaoyu Ji [1], Wenyuan Xu [1].
13
+
14
+ This repository is an official implementation of the SafeEar accepted to **ACM CCS 2024** (Core-A*, CCF-A, Big4) .
15
+
16
+ Please also visit our <a href="https://safeearweb.github.io/Project/">(1) Project Website</a>, <a href="https://zenodo.org/records/14062964">(2) Full CVoiceFake Dataset</a>, and <a href="https://zenodo.org/records/11124319">(3) Sampled CVoiceFake Dataset</a>.
17
+
18
+ ## 🔥News
19
+
20
+ [2025-03-18]: Supported the batch testing for ASVspoof 2019 and 2021, fixed some bugs for datasets and trainer.
21
+
22
+ [2024-12-10]: Fixed all the bugs for training and test, and uploaded the files for data generation `datas/`.
23
+
24
+ [2024-12-01]: Uploaded the checkpoint for data generation `datas/`.
25
+
26
+ ## ✨Key Highlights:
27
+
28
+ In this paper, we propose SafeEar, a novel framework that aims to detect deepfake audios without relying on accessing the speech content within. Our key idea is to devise a neural audio codec into a novel decoupling model that well separates the semantic and acoustic information from audio samples, and only use the acoustic information (e.g., prosody and timbre) for deepfake detection. In this way, no semantic content will be exposed to the detector. To overcome the challenge of identifying diverse deepfake audio without semantic clues, we enhance our deepfake detector with multi-head self-attention and codec augmentation. Extensive experiments conducted on four benchmark datasets demonstrate SafeEar’s effectiveness in detecting various deepfake techniques with an equal error rate (EER) down to 2.02%. Simultaneously, it shields five-language speech content from being deciphered by both machine and human auditory analysis, demonstrated by word error rates (WERs) all above 93.93% and our user study. Furthermore, our benchmark constructed for anti-deepfake and anti-content recovery evaluation helps provide a basis for future research in the realms of audio privacy preservation and deepfake detection.
29
+
30
+ ## 🚀Overall Pipeline
31
+
32
+ ![pipeline](assert/overall.gif)
33
+
34
+ ## 🔧Installation
35
+
36
+ 1. Clone the repository:
37
+
38
+ ```shell
39
+ git clone git@github.com:LetterLiGo/SafeEar.git
40
+ cd SafeEar/
41
+ ```
42
+
43
+ 2. Create and activate the conda environment:
44
+
45
+ ```shell
46
+ conda create -n safeear python=3.9
47
+ conda activate safeear
48
+ ```
49
+
50
+ 3. Install PyTorch and torchvision following the [official instructions](https://pytorch.org). The code requires `python=3.9`, `pytorch=1.13`, `torchvision=0.14`.
51
+
52
+
53
+ ```shell
54
+ pip install torch==1.13.1+cu116 torchvision==0.14.1+cu116 torchaudio==0.13.1 --extra-index-url https://download.pytorch.org/whl/cu116
55
+
56
+ ```
57
+ 4. Install other dependencies:
58
+
59
+ ```shell
60
+ pip install pip==24.0
61
+ pip install -r requirements.txt
62
+ ```
63
+
64
+ ## 📊Model Performance
65
+ ### ASVspoof 2019 & 2021
66
+ ![](assert/ASVSpoof-results.png)
67
+ ### Speech Recognition Performance
68
+ ![](assert/exp1.png)
69
+
70
+ ## Data preparation
71
+
72
+ ### AVSpoof 2019 & 2021
73
+
74
+ Please download the [ASVspoof 2019](https://datashare.is.ed.ac.uk/handle/10283/3336) and [ASVspoof 2021](https://www.asvspoof.org/index2021.html) datasets and extract them to the `datas/datasets` directory.
75
+
76
+ ```shell
77
+ datas/datasets/ASVspoof2019
78
+ datas/datasets/ASVspoof2021
79
+ ```
80
+
81
+ #### Generate the Hubert L9 feature files
82
+
83
+ ```shell
84
+ mkdir model_zoos
85
+ cd model_zoos
86
+ wget https://dl.fbaipublicfiles.com/hubert/hubert_base_ls960.pt
87
+ wget https://cloud.tsinghua.edu.cn/f/413a0cd2e6f749eea956/?dl=1 -O SpeechTokenizer.pt
88
+ cd ../datas
89
+ # Generate the Hubert L9 feature files for ASVspoof 2019
90
+ python dump_hubert_avg_feature.py datasets/ASVSpoof2019 datasets/ASVSpoof2019_Hubert_L9
91
+ # Generate the Hubert L9 feature files for ASVspoof 2021
92
+ python dump_hubert_avg_feature.py datasets/ASVSpoof2021 datasets/ASVSpoof2021_Hubert_L9
93
+ ```
94
+
95
+ ## 📚Training
96
+
97
+ Before starting training, please modify the parameter configurations in [`configs`](configs).
98
+
99
+ Use the following commands to start training:
100
+
101
+ ```shell
102
+ python train.py --conf_dir config/train19.yaml
103
+ python train.py --conf_dir config/train21.yaml
104
+ ```
105
+
106
+ ## 📈Testing/Inference
107
+
108
+ To evaluate a model on one or more GPUs, specify the `CUDA_VISIBLE_DEVICES`, `dataset`, `model` and `checkpoint`:
109
+
110
+ ```shell
111
+ python test.py --conf_dir Exps/ASVspoof19/config.yaml
112
+ python test.py --conf_dir Exps/ASVspoof21/config.yaml
113
+ ```
114
+
115
+ ## Bugs and Issues
116
+
117
+ If you meet `RuntimeError: Failed to load audio from <_io.BytesIO object at 0x7f45cb978f90>`, please use the following command to fix it:
118
+
119
+ ```shell
120
+ conda install -c anaconda 'ffmpeg<4.4'
121
+ ```
122
+
123
+ ## 📜Citation
124
+
125
+ If you find our work/code/dataset helpful, please consider citing:
126
+
127
+ ```
128
+ @inproceedings{li2024safeear,
129
+ author = {Li, Xinfeng and Li, Kai and Zheng, Yifan and Yan, Chen and Ji, Xiaoyu, and Xu, Wenyuan},
130
+ title = {{SafeEar: Content Privacy-Preserving Audio Deepfake Detection}},
131
+ booktitle = {Proceedings of the 2024 {ACM} {SIGSAC} Conference on Computer and Communications Security (CCS)}
132
+ year = {2024},
133
+ }
134
+ ```
clean/audio/safeear/SOURCE.md ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Source: audio/safeear
2
+
3
+ | Field | Value |
4
+ |---|---|
5
+ | Upstream | **UNVERIFIED** -- provenance was lost when this code was vendored |
6
+ | Paper | https://arxiv.org/abs/2409.09272 |
7
+ | Commit SHA | **not recorded** -- the vendoring step did not preserve it |
8
+ | Mirrored on | 2026-09-15 |
9
+ | Upstream license | LICENSE |
10
+
11
+ This is a **mirror**, stripped to the files needed for inference. The full
12
+ untouched snapshot is at `archive/audio__safeear.tar.gz`.
13
+
14
+ This code is the work of its original authors and is **not** covered by the
15
+ DeepSafe project license. If you are an author and want this removed, open an
16
+ issue on https://github.com/deepsafehq/deepsafe-bench and it will be taken down
17
+ within 48 hours, no questions asked.
clean/audio/safeear/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
clean/audio/safeear/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
clean/audio/safeear/requirements.txt ADDED
@@ -0,0 +1,123 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ absl-py==2.1.0
2
+ aiohttp==3.9.0
3
+ aiosignal==1.3.1
4
+ antlr4-python3-runtime==4.8
5
+ appdirs==1.4.4
6
+ asttokens==2.4.1
7
+ async-timeout==4.0.3
8
+ attrs==23.1.0
9
+ audioread==3.0.1
10
+ bitarray==2.8.3
11
+ blessed==1.20.0
12
+ certifi==2022.12.7
13
+ cffi==1.16.0
14
+ charset-normalizer==2.1.1
15
+ click==8.1.7
16
+ cmake==3.25.0
17
+ colorama==0.4.6
18
+ contourpy==1.2.0
19
+ cycler==0.12.1
20
+ Cython==3.0.5
21
+ decorator==5.1.1
22
+ docker-pycreds==0.4.0
23
+ einops==0.7.0
24
+ exceptiongroup==1.2.0
25
+ executing==2.0.1
26
+ # Editable install with no version control (fairseq==1.0.0a0)
27
+ -e fairseq_ours
28
+ fast-bss-eval==0.1.4
29
+ filelock==3.9.0
30
+ fonttools==4.45.0
31
+ frozenlist==1.4.0
32
+ fsspec==2023.10.0
33
+ gitdb==4.0.11
34
+ GitPython==3.1.40
35
+ gpustat==1.1.1
36
+ grpcio==1.63.0
37
+ huggingface-hub==0.19.4
38
+ hydra-core==1.0.7
39
+ idna==3.4
40
+ importlib-resources==6.1.1
41
+ importlib_metadata==7.1.0
42
+ ipdb==0.13.13
43
+ ipython==8.18.1
44
+ jedi==0.19.1
45
+ Jinja2==3.1.2
46
+ joblib==1.4.2
47
+ kiwisolver==1.4.5
48
+ lazy_loader==0.4
49
+ librosa==0.10.2
50
+ lightning-utilities==0.10.0
51
+ lit==15.0.7
52
+ llvmlite==0.42.0
53
+ lxml==4.9.3
54
+ Markdown==3.6
55
+ markdown-it-py==3.0.0
56
+ MarkupSafe==2.1.3
57
+ matplotlib==3.8.2
58
+ matplotlib-inline==0.1.6
59
+ mdurl==0.1.2
60
+ mpmath==1.3.0
61
+ msgpack==1.0.8
62
+ multidict==6.0.4
63
+ networkx==3.0
64
+ numba==0.59.1
65
+ numpy==1.23.5
66
+ nvidia-ml-py==12.535.133
67
+ opencv-python==4.9.0.80
68
+ packaging==23.2
69
+ parso==0.8.3
70
+ pexpect==4.9.0
71
+ Pillow==9.3.0
72
+ platformdirs==4.2.1
73
+ pooch==1.8.1
74
+ portalocker==2.8.2
75
+ prompt-toolkit==3.0.43
76
+ protobuf==4.25.1
77
+ psutil==5.9.6
78
+ ptyprocess==0.7.0
79
+ pure-eval==0.2.2
80
+ pycparser==2.21
81
+ pyDeprecate==0.3.2
82
+ Pygments==2.17.2
83
+ pyparsing==3.1.1
84
+ python-dateutil==2.8.2
85
+ pytorch-lightning==1.6.3
86
+ pytorch-ranger==0.1.1
87
+ PyYAML==6.0.1
88
+ regex==2023.10.3
89
+ requests==2.28.1
90
+ rich==13.7.0
91
+ sacrebleu==2.3.2
92
+ safetensors==0.4.0
93
+ scikit-learn==1.4.2
94
+ scipy==1.11.4
95
+ sentry-sdk==1.36.0
96
+ setproctitle==1.3.3
97
+ six==1.16.0
98
+ smmap==5.0.1
99
+ soundfile>=0.11.0
100
+ soxr==0.3.7
101
+ stack-data==0.6.3
102
+ sympy==1.12
103
+ tabulate==0.9.0
104
+ tensorboard==2.16.2
105
+ tensorboard-data-server==0.7.2
106
+ thop==0.1.1.post2209072238
107
+ threadpoolctl==3.5.0
108
+ timm==0.9.11
109
+ tomli==2.0.1
110
+ torch-mir-eval==0.4
111
+ torch-optimizer==0.3.0
112
+ torchmetrics==1.2.0
113
+ tqdm==4.66.1
114
+ traitlets==5.14.1
115
+ triton==2.0.0
116
+ typing_extensions==4.4.0
117
+ urllib3==1.26.13
118
+ wandb==0.16.0
119
+ wcwidth==0.2.12
120
+ Werkzeug==3.0.2
121
+ yarl==1.9.3
122
+ zipp==3.17.0
123
+ npy_append_array==0.9.16
clean/audio/safeear/safeear/losses/loss.py ADDED
@@ -0,0 +1,215 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn.functional as F
3
+ from torchaudio.transforms import MelSpectrogram
4
+ import numpy as np
5
+
6
+ def adversarial_g_loss(y_disc_gen):
7
+ """Hinge loss"""
8
+ loss = 0.0
9
+ for i in range(len(y_disc_gen)):
10
+ stft_loss = F.relu(1 - y_disc_gen[i]).mean().squeeze()
11
+ loss += stft_loss
12
+ return loss / len(y_disc_gen)
13
+
14
+
15
+ def feature_loss(fmap_r, fmap_gen):
16
+ loss = 0.0
17
+ for i in range(len(fmap_r)):
18
+ for j in range(len(fmap_r[i])):
19
+ stft_loss = ((fmap_r[i][j] - fmap_gen[i][j]).abs() /
20
+ (fmap_r[i][j].abs().mean())).mean()
21
+ loss += stft_loss
22
+ return loss / (len(fmap_r) * len(fmap_r[0]))
23
+
24
+
25
+ def sim_loss(y_disc_r, y_disc_gen):
26
+ loss = 0.0
27
+ for i in range(len(y_disc_r)):
28
+ loss += F.mse_loss(y_disc_r[i], y_disc_gen[i])
29
+ return loss / len(y_disc_r)
30
+
31
+ def reconstruction_loss(x, G_x, lamdba_wav=100, sr=16000, eps=1e-7):
32
+ # NOTE (lsx): hard-coded now
33
+ L = lamdba_wav * F.mse_loss(x, G_x) # wav L1 loss
34
+ # loss_sisnr = sisnr_loss(G_x, x) #
35
+ # L += 0.01*loss_sisnr
36
+ # 2^6=64 -> 2^10=1024
37
+ # NOTE (lsx): add 2^11
38
+ for i in range(6, 12):
39
+ # for i in range(5, 12): # Encodec setting
40
+ s = 2**i
41
+ melspec = MelSpectrogram(
42
+ sample_rate=sr,
43
+ n_fft=s,
44
+ hop_length=s // 4,
45
+ n_mels=64,
46
+ wkwargs={"device": x.device}).to(x.device)
47
+ S_x = melspec(x)
48
+ S_G_x = melspec(G_x)
49
+ loss = ((S_x - S_G_x).abs().mean() + (
50
+ ((torch.log(S_x.abs() + eps) - torch.log(S_G_x.abs() + eps))**2
51
+ ).mean(dim=-2)**0.5).mean()) / i
52
+ L += loss
53
+ return L
54
+
55
+
56
+ def criterion_d(y_disc_r, y_disc_gen, fmap_r_det, fmap_gen_det, y_df_hat_r,
57
+ y_df_hat_g, fmap_f_r, fmap_f_g, y_ds_hat_r, y_ds_hat_g,
58
+ fmap_s_r, fmap_s_g):
59
+ """Hinge Loss"""
60
+ loss = 0.0
61
+ loss1 = 0.0
62
+ loss2 = 0.0
63
+ loss3 = 0.0
64
+ for i in range(len(y_disc_r)):
65
+ loss1 += F.relu(1 - y_disc_r[i]).mean() + F.relu(1 + y_disc_gen[
66
+ i]).mean()
67
+ for i in range(len(y_df_hat_r)):
68
+ loss2 += F.relu(1 - y_df_hat_r[i]).mean() + F.relu(1 + y_df_hat_g[
69
+ i]).mean()
70
+ for i in range(len(y_ds_hat_r)):
71
+ loss3 += F.relu(1 - y_ds_hat_r[i]).mean() + F.relu(1 + y_ds_hat_g[
72
+ i]).mean()
73
+
74
+ loss = (loss1 / len(y_disc_gen) + loss2 / len(y_df_hat_r) + loss3 /
75
+ len(y_ds_hat_r)) / 3.0
76
+
77
+ return loss
78
+
79
+
80
+ def criterion_g(commit_loss, x, G_x, fmap_r, fmap_gen, y_disc_r, y_disc_gen,
81
+ y_df_hat_r, y_df_hat_g, fmap_f_r, fmap_f_g, y_ds_hat_r,
82
+ y_ds_hat_g, fmap_s_r, fmap_s_g, lamdba_wav=100, lamdba_com=1000, lamdba_adv=1, lamdba_feat=1, lamdba_rec=1, sr=16000):
83
+ adv_g_loss = adversarial_g_loss(y_disc_gen)
84
+ feat_loss = (feature_loss(fmap_r, fmap_gen) + sim_loss(
85
+ y_disc_r, y_disc_gen) + feature_loss(fmap_f_r, fmap_f_g) + sim_loss(
86
+ y_df_hat_r, y_df_hat_g) + feature_loss(fmap_s_r, fmap_s_g) +
87
+ sim_loss(y_ds_hat_r, y_ds_hat_g)) / 3.0
88
+ rec_loss = reconstruction_loss(x.contiguous(), G_x.contiguous(), lamdba_wav, sr)
89
+ total_loss = lamdba_com * commit_loss + lamdba_adv * adv_g_loss + lamdba_feat * feat_loss + lamdba_rec * rec_loss
90
+ return total_loss, adv_g_loss, feat_loss, rec_loss
91
+
92
+
93
+ def adopt_weight(weight, global_step, threshold=0, value=0.):
94
+ if global_step < threshold:
95
+ weight = value
96
+ return weight
97
+
98
+
99
+ def adopt_dis_weight(weight, global_step, threshold=0, value=0.):
100
+ # 0,3,6,9,13....这些时间步,不更新dis
101
+ if global_step % 3 == 0:
102
+ weight = value
103
+ return weight
104
+
105
+
106
+ def calculate_adaptive_weight(nll_loss, g_loss, last_layer, lamdba_adv=1):
107
+ if last_layer is not None:
108
+ nll_grads = torch.autograd.grad(
109
+ nll_loss, last_layer, retain_graph=True)[0]
110
+ g_grads = torch.autograd.grad(g_loss, last_layer, retain_graph=True)[0]
111
+ else:
112
+ print('last_layer cannot be none')
113
+ assert 1 == 2
114
+ d_weight = torch.norm(nll_grads) / (torch.norm(g_grads) + 1e-4)
115
+ d_weight = torch.clamp(d_weight, 1.0, 1.0).detach()
116
+ d_weight = d_weight * lamdba_adv
117
+ return d_weight
118
+
119
+ def loss_g(codebook_loss,
120
+ inputs,
121
+ reconstructions,
122
+ fmap_r,
123
+ fmap_gen,
124
+ y_disc_r,
125
+ y_disc_gen,
126
+ global_step,
127
+ y_df_hat_r,
128
+ y_df_hat_g,
129
+ y_ds_hat_r,
130
+ y_ds_hat_g,
131
+ fmap_f_r,
132
+ fmap_f_g,
133
+ fmap_s_r,
134
+ fmap_s_g,
135
+ lamdba_wav=100,
136
+ lamdba_com=1000,
137
+ lamdba_adv=1,
138
+ lamdba_feat=1,
139
+ sr=16000,
140
+ discriminator_iter_start=500
141
+ ):
142
+ """
143
+ args:
144
+ codebook_loss: commit loss.
145
+ inputs: ground-truth wav.
146
+ reconstructions: reconstructed wav.
147
+ fmap_r: real stft-D feature map.
148
+ fmap_gen: fake stft-D feature map.
149
+ y_disc_r: real stft-D logits.
150
+ y_disc_gen: fake stft-D logits.
151
+ global_step: global training step.
152
+ y_df_hat_r: real MPD logits.
153
+ y_df_hat_g: fake MPD logits.
154
+ y_ds_hat_r: real MSD logits.
155
+ y_ds_hat_g: fake MSD logits.
156
+ fmap_f_r: real MPD feature map.
157
+ fmap_f_g: fake MPD feature map.
158
+ fmap_s_r: real MSD feature map.
159
+ fmap_s_g: fake MSD feature map.
160
+ """
161
+ rec_loss = reconstruction_loss(inputs.contiguous(),
162
+ reconstructions.contiguous(), lamdba_wav, sr)
163
+ adv_g_loss = adversarial_g_loss(y_disc_gen)
164
+ adv_mpd_loss = adversarial_g_loss(y_df_hat_g)
165
+ adv_msd_loss = adversarial_g_loss(y_ds_hat_g)
166
+ adv_loss = (adv_g_loss + adv_mpd_loss + adv_msd_loss
167
+ ) / 3.0 # NOTE(lsx): need to divide by 3?
168
+ feat_loss = feature_loss(
169
+ fmap_r,
170
+ fmap_gen) #+ sim_loss(y_disc_r, y_disc_gen) # NOTE(lsx): need logits?
171
+ feat_loss_mpd = feature_loss(fmap_f_r,
172
+ fmap_f_g) #+ sim_loss(y_df_hat_r, y_df_hat_g)
173
+ feat_loss_msd = feature_loss(fmap_s_r,
174
+ fmap_s_g) #+ sim_loss(y_ds_hat_r, y_ds_hat_g)
175
+ feat_loss_tot = (feat_loss + feat_loss_mpd + feat_loss_msd) / 3.0
176
+ d_weight = torch.tensor(1.0)
177
+ disc_factor = adopt_weight(
178
+ lamdba_adv, global_step, threshold=discriminator_iter_start)
179
+ if disc_factor == 0.:
180
+ fm_loss_wt = 0
181
+ else:
182
+ fm_loss_wt = lamdba_feat
183
+ loss = rec_loss + d_weight * disc_factor * adv_loss + \
184
+ fm_loss_wt * feat_loss_tot + lamdba_com * codebook_loss
185
+ return loss, rec_loss, adv_loss, feat_loss_tot, d_weight
186
+
187
+ def compute_det_curve(target_scores, nontarget_scores):
188
+
189
+ n_scores = target_scores.size + nontarget_scores.size
190
+ all_scores = np.concatenate((target_scores, nontarget_scores))
191
+ labels = np.concatenate((np.ones(target_scores.size), np.zeros(nontarget_scores.size)))
192
+
193
+ # Sort labels based on scores
194
+ indices = np.argsort(all_scores, kind='mergesort')
195
+ labels = labels[indices]
196
+
197
+ # Compute false rejection and false acceptance rates
198
+ tar_trial_sums = np.cumsum(labels)
199
+ nontarget_trial_sums = nontarget_scores.size - (np.arange(1, n_scores + 1) - tar_trial_sums)
200
+
201
+ frr = np.concatenate((np.atleast_1d(0), tar_trial_sums / target_scores.size)) # false rejection rates
202
+ far = np.concatenate((np.atleast_1d(1), nontarget_trial_sums / nontarget_scores.size)) # false acceptance rates
203
+ thresholds = np.concatenate((np.atleast_1d(all_scores[indices[0]] - 0.001), all_scores[indices])) # Thresholds are the sorted scores
204
+
205
+ return frr, far, thresholds
206
+
207
+
208
+ def compute_eer(target_scores, nontarget_scores):
209
+ """ Returns equal error rate (EER) and the corresponding threshold. """
210
+ frr, far, thresholds = compute_det_curve(target_scores, nontarget_scores)
211
+ abs_diffs = np.abs(frr - far)
212
+ min_index = np.argmin(abs_diffs)
213
+ eer = np.mean((frr[min_index], far[min_index]))
214
+ print(thresholds[min_index])
215
+ return eer, thresholds[min_index]
clean/audio/safeear/safeear/models/decouple.py ADDED
@@ -0,0 +1,207 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # -*- coding: utf-8 -*-
2
+ """
3
+ Created on Wed Aug 30 15:47:55 2023
4
+ @author: zhangxin
5
+ """
6
+ import torch.nn as nn
7
+ from einops import rearrange
8
+ import torch
9
+
10
+ from .modules.seanet import SEANetEncoder, SEANetDecoder
11
+ from .modules.quantization import ResidualVectorQuantizer
12
+
13
+
14
+ class SpeechTokenizer(nn.Module):
15
+ def __init__(self, n_filters, dimension, strides, lstm_layers, bidirectional, dilation_base, residual_kernel_size, n_residual_layers, activation, sample_rate, n_q, semantic_dimension, codebook_size):
16
+ '''
17
+
18
+ Parameters
19
+ ----------
20
+ n_filters : int
21
+ Number of filters in the SEANet encoder/decoder.
22
+ dimension : int
23
+ Dimensionality of the encoder/decoder.
24
+ strides : list
25
+ List of stride values for the SEANet encoder/decoder.
26
+ lstm_layers : int
27
+ Number of LSTM layers in the encoder/decoder.
28
+ bidirectional : bool
29
+ Whether to use bidirectional LSTM in the encoder.
30
+ dilation_base : int
31
+ Base dilation rate for the residual blocks in the encoder/decoder.
32
+ residual_kernel_size : int
33
+ Kernel size for the residual blocks in the encoder/decoder.
34
+ n_residual_layers : int
35
+ Number of residual layers in the encoder/decoder.
36
+ activation : str
37
+ Activation function to use in the encoder/decoder.
38
+ sample_rate : int
39
+ Sample rate of the audio.
40
+ n_q : int
41
+ Number of quantization levels.
42
+ semantic_dimension : int
43
+ Dimensionality of the semantic representation.
44
+ codebook_size : int
45
+ Size of the codebook for vector quantization.
46
+
47
+ '''
48
+ super().__init__()
49
+ self.encoder = SEANetEncoder(n_filters=n_filters,
50
+ dimension=dimension,
51
+ ratios=strides,
52
+ lstm=lstm_layers,
53
+ bidirectional=bidirectional,
54
+ dilation_base=dilation_base,
55
+ residual_kernel_size=residual_kernel_size,
56
+ n_residual_layers=n_residual_layers,
57
+ activation=activation)
58
+ self.sample_rate = sample_rate
59
+ self.n_q = n_q
60
+ if dimension != semantic_dimension:
61
+ self.transform = nn.Linear(dimension, semantic_dimension)
62
+ else:
63
+ self.transform = nn.Identity()
64
+ self.quantizer = ResidualVectorQuantizer(dimension=dimension, n_q=n_q, bins=codebook_size)
65
+ self.decoder = SEANetDecoder(n_filters=n_filters,
66
+ dimension=dimension,
67
+ ratios=strides,
68
+ lstm=lstm_layers,
69
+ bidirectional=False,
70
+ dilation_base=dilation_base,
71
+ residual_kernel_size=residual_kernel_size,
72
+ n_residual_layers=n_residual_layers,
73
+ activation=activation)
74
+
75
+ @classmethod
76
+ def load_from_checkpoint(cls,
77
+ config_path: str,
78
+ ckpt_path: str):
79
+ '''
80
+
81
+ Parameters
82
+ ----------
83
+ config_path : str
84
+ Path of model configuration file.
85
+ ckpt_path : str
86
+ Path of model checkpoint.
87
+
88
+ Returns
89
+ -------
90
+ model : SpeechTokenizer
91
+ SpeechTokenizer model.
92
+
93
+ '''
94
+ import json
95
+ with open(config_path) as f:
96
+ cfg = json.load(f)
97
+ model = cls(cfg)
98
+ params = torch.load(ckpt_path, map_location='cpu')
99
+ model.load_state_dict(params)
100
+ return model
101
+
102
+
103
+ def forward(self,
104
+ x: torch.tensor,
105
+ n_q: int=None,
106
+ layers: list=[0]):
107
+ '''
108
+
109
+ Parameters
110
+ ----------
111
+ x : torch.tensor
112
+ Input wavs. Shape: (batch, channels, timesteps).
113
+ n_q : int, optional
114
+ Number of quantizers in RVQ used to encode. The default is all layers.
115
+ layers : list[int], optional
116
+ Layers of RVQ should return quantized result. The default is the first layer.
117
+
118
+ Returns
119
+ -------
120
+ o : torch.tensor
121
+ Output wavs. Shape: (batch, channels, timesteps).
122
+ commit_loss : torch.tensor
123
+ Commitment loss from residual vector quantizers.
124
+ feature : torch.tensor
125
+ Output of RVQ's first layer. Shape: (batch, timesteps, dimension)
126
+
127
+ '''
128
+ n_q = n_q if n_q else self.n_q
129
+ e = self.encoder(x)
130
+ quantized, codes, commit_loss, quantized_list = self.quantizer(e, n_q=n_q, layers=layers)
131
+ feature = rearrange(quantized_list[0], 'b d t -> b t d') # b,t,1024
132
+ feature = self.transform(feature) #b,t,768
133
+ o = self.decoder(quantized)
134
+ return o, commit_loss, feature, quantized_list[1:]
135
+
136
+ def forward_feature(self,
137
+ x: torch.tensor,
138
+ layers: list=None):
139
+ '''
140
+
141
+ Parameters
142
+ ----------
143
+ x : torch.tensor
144
+ Input wavs. Shape should be (batch, channels, timesteps).
145
+ layers : list[int], optional
146
+ Layers of RVQ should return quantized result. The default is all layers.
147
+
148
+ Returns
149
+ -------
150
+ quantized_list : list[torch.tensor]
151
+ Quantized of required layers.
152
+
153
+ '''
154
+ e = self.encoder(x)
155
+ layers = layers if layers else list(range(self.n_q))
156
+ quantized, codes, commit_loss, quantized_list = self.quantizer(e, layers=layers)
157
+ return quantized_list
158
+
159
+ def encode(self,
160
+ x: torch.tensor,
161
+ n_q: int=None,
162
+ st: int=None):
163
+ '''
164
+
165
+ Parameters
166
+ ----------
167
+ x : torch.tensor
168
+ Input wavs. Shape: (batch, channels, timesteps).
169
+ n_q : int, optional
170
+ Number of quantizers in RVQ used to encode. The default is all layers.
171
+ st : int, optional
172
+ Start quantizer index in RVQ. The default is 0.
173
+
174
+ Returns
175
+ -------
176
+ codes : torch.tensor
177
+ Output indices for each quantizer. Shape: (n_q, batch, timesteps)
178
+
179
+ '''
180
+ e = self.encoder(x)
181
+ if st is None:
182
+ st = 0
183
+ n_q = n_q if n_q else self.n_q
184
+ codes = self.quantizer.encode(e, n_q=n_q, st=st)
185
+ return codes
186
+
187
+ def decode(self,
188
+ codes: torch.tensor,
189
+ st: int=0):
190
+ '''
191
+
192
+ Parameters
193
+ ----------
194
+ codes : torch.tensor
195
+ Indices for each quantizer. Shape: (n_q, batch, timesteps).
196
+ st : int, optional
197
+ Start quantizer index in RVQ. The default is 0.
198
+
199
+ Returns
200
+ -------
201
+ o : torch.tensor
202
+ Reconstruct wavs from codes. Shape: (batch, channels, timesteps)
203
+
204
+ '''
205
+ quantized = self.quantizer.decode(codes, st=st)
206
+ o = self.decoder(quantized)
207
+ return o
clean/audio/safeear/safeear/models/discriminator.py ADDED
@@ -0,0 +1,422 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+ """MS-STFT discriminator, provided here for reference."""
7
+ import typing as tp
8
+
9
+ import torch
10
+ import torchaudio
11
+ from einops import rearrange
12
+ from torch import nn
13
+ from torch.nn import functional as F
14
+ import einops
15
+ from torch.nn import AvgPool1d
16
+ from torch.nn.utils import spectral_norm
17
+ from torch.nn.utils import weight_norm
18
+
19
+ FeatureMapType = tp.List[torch.Tensor]
20
+ LogitsType = torch.Tensor
21
+ DiscriminatorOutput = tp.Tuple[tp.List[LogitsType], tp.List[FeatureMapType]]
22
+
23
+ CONV_NORMALIZATIONS = frozenset([
24
+ 'none', 'weight_norm', 'spectral_norm', 'time_layer_norm', 'layer_norm',
25
+ 'time_group_norm'
26
+ ])
27
+
28
+ class ConvLayerNorm(nn.LayerNorm):
29
+ """
30
+ Convolution-friendly LayerNorm that moves channels to last dimensions
31
+ before running the normalization and moves them back to original position right after.
32
+ """
33
+
34
+ def __init__(self,
35
+ normalized_shape: tp.Union[int, tp.List[int], torch.Size],
36
+ **kwargs):
37
+ super().__init__(normalized_shape, **kwargs)
38
+
39
+ def forward(self, x):
40
+ x = einops.rearrange(x, 'b ... t -> b t ...')
41
+ x = super().forward(x)
42
+ x = einops.rearrange(x, 'b t ... -> b ... t')
43
+ return
44
+
45
+
46
+ def apply_parametrization_norm(module: nn.Module,
47
+ norm: str='none') -> nn.Module:
48
+ assert norm in CONV_NORMALIZATIONS
49
+ if norm == 'weight_norm':
50
+ return weight_norm(module)
51
+ elif norm == 'spectral_norm':
52
+ return spectral_norm(module)
53
+ else:
54
+ # We already check was in CONV_NORMALIZATION, so any other choice
55
+ # doesn't need reparametrization.
56
+ return module
57
+
58
+
59
+ def get_norm_module(module: nn.Module,
60
+ causal: bool=False,
61
+ norm: str='none',
62
+ **norm_kwargs) -> nn.Module:
63
+ """Return the proper normalization module. If causal is True, this will ensure the returned
64
+ module is causal, or return an error if the normalization doesn't support causal evaluation.
65
+ """
66
+ assert norm in CONV_NORMALIZATIONS
67
+ if norm == 'layer_norm':
68
+ assert isinstance(module, nn.modules.conv._ConvNd)
69
+ return ConvLayerNorm(module.out_channels, **norm_kwargs)
70
+ elif norm == 'time_group_norm':
71
+ if causal:
72
+ raise ValueError("GroupNorm doesn't support causal evaluation.")
73
+ assert isinstance(module, nn.modules.conv._ConvNd)
74
+ return nn.GroupNorm(1, module.out_channels, **norm_kwargs)
75
+ else:
76
+ return nn.Identity()
77
+
78
+ def get_padding(kernel_size, dilation=1):
79
+ return int((kernel_size * dilation - dilation) / 2)
80
+
81
+ class NormConv1d(nn.Module):
82
+ """Wrapper around Conv1d and normalization applied to this conv
83
+ to provide a uniform interface across normalization approaches.
84
+ """
85
+
86
+ def __init__(self,
87
+ *args,
88
+ causal: bool=False,
89
+ norm: str='none',
90
+ norm_kwargs: tp.Dict[str, tp.Any]={},
91
+ **kwargs):
92
+ super().__init__()
93
+ self.conv = apply_parametrization_norm(nn.Conv1d(*args, **kwargs), norm)
94
+ self.norm = get_norm_module(self.conv, causal, norm, **norm_kwargs)
95
+ self.norm_type = norm
96
+
97
+ def forward(self, x):
98
+ x = self.conv(x)
99
+ x = self.norm(x)
100
+ return x
101
+
102
+
103
+ class NormConv2d(nn.Module):
104
+ """Wrapper around Conv2d and normalization applied to this conv
105
+ to provide a uniform interface across normalization approaches.
106
+ """
107
+
108
+ def __init__(self,
109
+ *args,
110
+ norm: str='none',
111
+ norm_kwargs: tp.Dict[str, tp.Any]={},
112
+ **kwargs):
113
+ super().__init__()
114
+ self.conv = apply_parametrization_norm(nn.Conv2d(*args, **kwargs), norm)
115
+ self.norm = get_norm_module(
116
+ self.conv, causal=False, norm=norm, **norm_kwargs)
117
+ self.norm_type = norm
118
+
119
+ def forward(self, x):
120
+ x = self.conv(x)
121
+ x = self.norm(x)
122
+ return x
123
+
124
+
125
+ def get_2d_padding(kernel_size: tp.Tuple[int, int],
126
+ dilation: tp.Tuple[int, int]=(1, 1)):
127
+ return (((kernel_size[0] - 1) * dilation[0]) // 2, (
128
+ (kernel_size[1] - 1) * dilation[1]) // 2)
129
+
130
+
131
+ class DiscriminatorSTFT(nn.Module):
132
+ """STFT sub-discriminator.
133
+ Args:
134
+ filters (int): Number of filters in convolutions
135
+ in_channels (int): Number of input channels. Default: 1
136
+ out_channels (int): Number of output channels. Default: 1
137
+ n_fft (int): Size of FFT for each scale. Default: 1024
138
+ hop_length (int): Length of hop between STFT windows for each scale. Default: 256
139
+ kernel_size (tuple of int): Inner Conv2d kernel sizes. Default: ``(3, 9)``
140
+ stride (tuple of int): Inner Conv2d strides. Default: ``(1, 2)``
141
+ dilations (list of int): Inner Conv2d dilation on the time dimension. Default: ``[1, 2, 4]``
142
+ win_length (int): Window size for each scale. Default: 1024
143
+ normalized (bool): Whether to normalize by magnitude after stft. Default: True
144
+ norm (str): Normalization method. Default: `'weight_norm'`
145
+ activation (str): Activation function. Default: `'LeakyReLU'`
146
+ activation_params (dict): Parameters to provide to the activation function.
147
+ growth (int): Growth factor for the filters. Default: 1
148
+ """
149
+
150
+ def __init__(self,
151
+ filters: int,
152
+ in_channels: int=1,
153
+ out_channels: int=1,
154
+ n_fft: int=1024,
155
+ hop_length: int=256,
156
+ win_length: int=1024,
157
+ max_filters: int=1024,
158
+ filters_scale: int=1,
159
+ kernel_size: tp.Tuple[int, int]=(3, 9),
160
+ dilations: tp.List=[1, 2, 4],
161
+ stride: tp.Tuple[int, int]=(1, 2),
162
+ normalized: bool=True,
163
+ norm: str='weight_norm',
164
+ activation: str='LeakyReLU',
165
+ activation_params: dict={'negative_slope': 0.2}):
166
+ super().__init__()
167
+ assert len(kernel_size) == 2
168
+ assert len(stride) == 2
169
+ self.filters = filters
170
+ self.in_channels = in_channels
171
+ self.out_channels = out_channels
172
+ self.n_fft = n_fft
173
+ self.hop_length = hop_length
174
+ self.win_length = win_length
175
+ self.normalized = normalized
176
+ self.activation = getattr(torch.nn, activation)(**activation_params)
177
+ self.spec_transform = torchaudio.transforms.Spectrogram(
178
+ n_fft=self.n_fft,
179
+ hop_length=self.hop_length,
180
+ win_length=self.win_length,
181
+ window_fn=torch.hann_window,
182
+ normalized=self.normalized,
183
+ center=False,
184
+ pad_mode=None,
185
+ power=None)
186
+ spec_channels = 2 * self.in_channels
187
+ self.convs = nn.ModuleList()
188
+ self.convs.append(
189
+ NormConv2d(
190
+ spec_channels,
191
+ self.filters,
192
+ kernel_size=kernel_size,
193
+ padding=get_2d_padding(kernel_size)))
194
+ in_chs = min(filters_scale * self.filters, max_filters)
195
+ for i, dilation in enumerate(dilations):
196
+ out_chs = min((filters_scale**(i + 1)) * self.filters, max_filters)
197
+ self.convs.append(
198
+ NormConv2d(
199
+ in_chs,
200
+ out_chs,
201
+ kernel_size=kernel_size,
202
+ stride=stride,
203
+ dilation=(dilation, 1),
204
+ padding=get_2d_padding(kernel_size, (dilation, 1)),
205
+ norm=norm))
206
+ in_chs = out_chs
207
+ out_chs = min((filters_scale**(len(dilations) + 1)) * self.filters,
208
+ max_filters)
209
+ self.convs.append(
210
+ NormConv2d(
211
+ in_chs,
212
+ out_chs,
213
+ kernel_size=(kernel_size[0], kernel_size[0]),
214
+ padding=get_2d_padding((kernel_size[0], kernel_size[0])),
215
+ norm=norm))
216
+ self.conv_post = NormConv2d(
217
+ out_chs,
218
+ self.out_channels,
219
+ kernel_size=(kernel_size[0], kernel_size[0]),
220
+ padding=get_2d_padding((kernel_size[0], kernel_size[0])),
221
+ norm=norm)
222
+
223
+ def forward(self, x: torch.Tensor):
224
+ fmap = []
225
+ # print('x ', x.shape)
226
+ z = self.spec_transform(x) # [B, 2, Freq, Frames, 2]
227
+ # print('z ', z.shape)
228
+ z = torch.cat([z.real, z.imag], dim=1)
229
+ # print('cat_z ', z.shape)
230
+ z = rearrange(z, 'b c w t -> b c t w')
231
+ for i, layer in enumerate(self.convs):
232
+ z = layer(z)
233
+ z = self.activation(z)
234
+ # print('z i', i, z.shape)
235
+ fmap.append(z)
236
+ z = self.conv_post(z)
237
+ # print('logit ', z.shape)
238
+ return z, fmap
239
+
240
+
241
+ class MultiScaleSTFTDiscriminator(nn.Module):
242
+ """Multi-Scale STFT (MS-STFT) discriminator.
243
+ Args:
244
+ filters (int): Number of filters in convolutions
245
+ in_channels (int): Number of input channels. Default: 1
246
+ out_channels (int): Number of output channels. Default: 1
247
+ n_ffts (Sequence[int]): Size of FFT for each scale
248
+ hop_lengths (Sequence[int]): Length of hop between STFT windows for each scale
249
+ win_lengths (Sequence[int]): Window size for each scale
250
+ **kwargs: additional args for STFTDiscriminator
251
+ """
252
+
253
+ def __init__(self,
254
+ filters: int,
255
+ in_channels: int=1,
256
+ out_channels: int=1,
257
+ n_ffts: tp.List[int]=[1024, 2048, 512, 256, 128],
258
+ hop_lengths: tp.List[int]=[256, 512, 128, 64, 32],
259
+ win_lengths: tp.List[int]=[1024, 2048, 512, 256, 128],
260
+ **kwargs):
261
+ super().__init__()
262
+ assert len(n_ffts) == len(hop_lengths) == len(win_lengths)
263
+ self.discriminators = nn.ModuleList([
264
+ DiscriminatorSTFT(
265
+ filters,
266
+ in_channels=in_channels,
267
+ out_channels=out_channels,
268
+ n_fft=n_ffts[i],
269
+ win_length=win_lengths[i],
270
+ hop_length=hop_lengths[i],
271
+ **kwargs) for i in range(len(n_ffts))
272
+ ])
273
+ self.num_discriminators = len(self.discriminators)
274
+
275
+ def forward(self, x: torch.Tensor) -> DiscriminatorOutput:
276
+ logits = []
277
+ fmaps = []
278
+ for disc in self.discriminators:
279
+ logit, fmap = disc(x)
280
+ logits.append(logit)
281
+ fmaps.append(fmap)
282
+ return logits, fmaps
283
+
284
+
285
+ class DiscriminatorP(torch.nn.Module):
286
+ def __init__(self,
287
+ period,
288
+ kernel_size=5,
289
+ stride=3,
290
+ use_spectral_norm=False,
291
+ activation: str='LeakyReLU',
292
+ activation_params: dict={'negative_slope': 0.2}):
293
+ super(DiscriminatorP, self).__init__()
294
+ self.period = period
295
+ norm_f = weight_norm if use_spectral_norm is False else spectral_norm
296
+ self.activation = getattr(torch.nn, activation)(**activation_params)
297
+ self.convs = nn.ModuleList([
298
+ NormConv2d(
299
+ 1,
300
+ 32, (kernel_size, 1), (stride, 1),
301
+ padding=(get_padding(5, 1), 0)),
302
+ NormConv2d(
303
+ 32,
304
+ 32, (kernel_size, 1), (stride, 1),
305
+ padding=(get_padding(5, 1), 0)),
306
+ NormConv2d(
307
+ 32,
308
+ 32, (kernel_size, 1), (stride, 1),
309
+ padding=(get_padding(5, 1), 0)),
310
+ NormConv2d(
311
+ 32,
312
+ 32, (kernel_size, 1), (stride, 1),
313
+ padding=(get_padding(5, 1), 0)),
314
+ NormConv2d(32, 32, (kernel_size, 1), 1, padding=(2, 0)),
315
+ ])
316
+ self.conv_post = NormConv2d(32, 1, (3, 1), 1, padding=(1, 0))
317
+
318
+ def forward(self, x):
319
+ fmap = []
320
+ # 1d to 2d
321
+ b, c, t = x.shape
322
+ if t % self.period != 0: # pad first
323
+ n_pad = self.period - (t % self.period)
324
+ x = F.pad(x, (0, n_pad), "reflect")
325
+ t = t + n_pad
326
+ x = x.view(b, c, t // self.period, self.period)
327
+
328
+ for l in self.convs:
329
+ x = l(x)
330
+ x = self.activation(x)
331
+ fmap.append(x)
332
+ x = self.conv_post(x)
333
+ fmap.append(x)
334
+ x = torch.flatten(x, 1, -1)
335
+
336
+ return x, fmap
337
+
338
+
339
+ class MultiPeriodDiscriminator(torch.nn.Module):
340
+ def __init__(self):
341
+ super(MultiPeriodDiscriminator, self).__init__()
342
+ self.discriminators = nn.ModuleList([
343
+ DiscriminatorP(2),
344
+ DiscriminatorP(3),
345
+ DiscriminatorP(5),
346
+ DiscriminatorP(7),
347
+ DiscriminatorP(11),
348
+ ])
349
+
350
+ def forward(self, y, y_hat):
351
+ y_d_rs = []
352
+ y_d_gs = []
353
+ fmap_rs = []
354
+ fmap_gs = []
355
+ for i, d in enumerate(self.discriminators):
356
+ y_d_r, fmap_r = d(y)
357
+ y_d_g, fmap_g = d(y_hat)
358
+ y_d_rs.append(y_d_r)
359
+ fmap_rs.append(fmap_r)
360
+ y_d_gs.append(y_d_g)
361
+ fmap_gs.append(fmap_g)
362
+ return y_d_rs, y_d_gs, fmap_rs, fmap_gs
363
+
364
+
365
+ class DiscriminatorS(torch.nn.Module):
366
+ def __init__(self,
367
+ use_spectral_norm=False,
368
+ activation: str='LeakyReLU',
369
+ activation_params: dict={'negative_slope': 0.2}):
370
+ super(DiscriminatorS, self).__init__()
371
+ self.activation = getattr(torch.nn, activation)(**activation_params)
372
+ self.convs = nn.ModuleList([
373
+ NormConv1d(1, 32, 15, 1, padding=7),
374
+ NormConv1d(32, 32, 41, 2, groups=4, padding=20),
375
+ NormConv1d(32, 32, 41, 2, groups=16, padding=20),
376
+ NormConv1d(32, 32, 41, 4, groups=16, padding=20),
377
+ NormConv1d(32, 32, 41, 4, groups=16, padding=20),
378
+ NormConv1d(32, 32, 41, 1, groups=16, padding=20),
379
+ NormConv1d(32, 32, 5, 1, padding=2),
380
+ ])
381
+ self.conv_post = NormConv1d(32, 1, 3, 1, padding=1)
382
+
383
+ def forward(self, x):
384
+ fmap = []
385
+ for l in self.convs:
386
+ x = l(x)
387
+ x = self.activation(x)
388
+ fmap.append(x)
389
+ x = self.conv_post(x)
390
+ fmap.append(x)
391
+ x = torch.flatten(x, 1, -1)
392
+ return x, fmap
393
+
394
+
395
+ class MultiScaleDiscriminator(torch.nn.Module):
396
+ def __init__(self):
397
+ super(MultiScaleDiscriminator, self).__init__()
398
+ self.discriminators = nn.ModuleList([
399
+ DiscriminatorS(),
400
+ DiscriminatorS(),
401
+ DiscriminatorS(),
402
+ ])
403
+ self.meanpools = nn.ModuleList(
404
+ [AvgPool1d(4, 2, padding=2), AvgPool1d(4, 2, padding=2)])
405
+
406
+ def forward(self, y, y_hat):
407
+ y_d_rs = []
408
+ y_d_gs = []
409
+ fmap_rs = []
410
+ fmap_gs = []
411
+ for i, d in enumerate(self.discriminators):
412
+ if i != 0:
413
+ y = self.meanpools[i - 1](y)
414
+ y_hat = self.meanpools[i - 1](y_hat)
415
+ y_d_r, fmap_r = d(y)
416
+ y_d_g, fmap_g = d(y_hat)
417
+ y_d_rs.append(y_d_r)
418
+ fmap_rs.append(fmap_r)
419
+ y_d_gs.append(y_d_g)
420
+ fmap_gs.append(fmap_g)
421
+
422
+ return y_d_rs, y_d_gs, fmap_rs, fmap_gs
clean/audio/safeear/safeear/models/modules/__init__.py ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ """Torch modules."""
8
+
9
+ # flake8: noqa
10
+ from .conv import (
11
+ pad1d,
12
+ unpad1d,
13
+ NormConv1d,
14
+ NormConvTranspose1d,
15
+ NormConv2d,
16
+ NormConvTranspose2d,
17
+ SConv1d,
18
+ SConvTranspose1d,
19
+ )
20
+ from .lstm import SLSTM
21
+ from .seanet import SEANetEncoder, SEANetDecoder
clean/audio/safeear/safeear/models/modules/conv.py ADDED
@@ -0,0 +1,252 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ """Convolutional layers wrappers and utilities."""
8
+
9
+ import math
10
+ import typing as tp
11
+ import warnings
12
+
13
+ import torch
14
+ from torch import nn
15
+ from torch.nn import functional as F
16
+ from torch.nn.utils import spectral_norm, weight_norm
17
+
18
+ from .norm import ConvLayerNorm
19
+
20
+
21
+ CONV_NORMALIZATIONS = frozenset(['none', 'weight_norm', 'spectral_norm',
22
+ 'time_layer_norm', 'layer_norm', 'time_group_norm'])
23
+
24
+
25
+ def apply_parametrization_norm(module: nn.Module, norm: str = 'none') -> nn.Module:
26
+ assert norm in CONV_NORMALIZATIONS
27
+ if norm == 'weight_norm':
28
+ return weight_norm(module)
29
+ elif norm == 'spectral_norm':
30
+ return spectral_norm(module)
31
+ else:
32
+ # We already check was in CONV_NORMALIZATION, so any other choice
33
+ # doesn't need reparametrization.
34
+ return module
35
+
36
+
37
+ def get_norm_module(module: nn.Module, causal: bool = False, norm: str = 'none', **norm_kwargs) -> nn.Module:
38
+ """Return the proper normalization module. If causal is True, this will ensure the returned
39
+ module is causal, or return an error if the normalization doesn't support causal evaluation.
40
+ """
41
+ assert norm in CONV_NORMALIZATIONS
42
+ if norm == 'layer_norm':
43
+ assert isinstance(module, nn.modules.conv._ConvNd)
44
+ return ConvLayerNorm(module.out_channels, **norm_kwargs)
45
+ elif norm == 'time_group_norm':
46
+ if causal:
47
+ raise ValueError("GroupNorm doesn't support causal evaluation.")
48
+ assert isinstance(module, nn.modules.conv._ConvNd)
49
+ return nn.GroupNorm(1, module.out_channels, **norm_kwargs)
50
+ else:
51
+ return nn.Identity()
52
+
53
+
54
+ def get_extra_padding_for_conv1d(x: torch.Tensor, kernel_size: int, stride: int,
55
+ padding_total: int = 0) -> int:
56
+ """See `pad_for_conv1d`.
57
+ """
58
+ length = x.shape[-1]
59
+ n_frames = (length - kernel_size + padding_total) / stride + 1
60
+ ideal_length = (math.ceil(n_frames) - 1) * stride + (kernel_size - padding_total)
61
+ return ideal_length - length
62
+
63
+
64
+ def pad_for_conv1d(x: torch.Tensor, kernel_size: int, stride: int, padding_total: int = 0):
65
+ """Pad for a convolution to make sure that the last window is full.
66
+ Extra padding is added at the end. This is required to ensure that we can rebuild
67
+ an output of the same length, as otherwise, even with padding, some time steps
68
+ might get removed.
69
+ For instance, with total padding = 4, kernel size = 4, stride = 2:
70
+ 0 0 1 2 3 4 5 0 0 # (0s are padding)
71
+ 1 2 3 # (output frames of a convolution, last 0 is never used)
72
+ 0 0 1 2 3 4 5 0 # (output of tr. conv., but pos. 5 is going to get removed as padding)
73
+ 1 2 3 4 # once you removed padding, we are missing one time step !
74
+ """
75
+ extra_padding = get_extra_padding_for_conv1d(x, kernel_size, stride, padding_total)
76
+ return F.pad(x, (0, extra_padding))
77
+
78
+
79
+ def pad1d(x: torch.Tensor, paddings: tp.Tuple[int, int], mode: str = 'zero', value: float = 0.):
80
+ """Tiny wrapper around F.pad, just to allow for reflect padding on small input.
81
+ If this is the case, we insert extra 0 padding to the right before the reflection happen.
82
+ """
83
+ length = x.shape[-1]
84
+ padding_left, padding_right = paddings
85
+ assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)
86
+ if mode == 'reflect':
87
+ max_pad = max(padding_left, padding_right)
88
+ extra_pad = 0
89
+ if length <= max_pad:
90
+ extra_pad = max_pad - length + 1
91
+ x = F.pad(x, (0, extra_pad))
92
+ padded = F.pad(x, paddings, mode, value)
93
+ end = padded.shape[-1] - extra_pad
94
+ return padded[..., :end]
95
+ else:
96
+ return F.pad(x, paddings, mode, value)
97
+
98
+
99
+ def unpad1d(x: torch.Tensor, paddings: tp.Tuple[int, int]):
100
+ """Remove padding from x, handling properly zero padding. Only for 1d!"""
101
+ padding_left, padding_right = paddings
102
+ assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)
103
+ assert (padding_left + padding_right) <= x.shape[-1]
104
+ end = x.shape[-1] - padding_right
105
+ return x[..., padding_left: end]
106
+
107
+
108
+ class NormConv1d(nn.Module):
109
+ """Wrapper around Conv1d and normalization applied to this conv
110
+ to provide a uniform interface across normalization approaches.
111
+ """
112
+ def __init__(self, *args, causal: bool = False, norm: str = 'none',
113
+ norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):
114
+ super().__init__()
115
+ self.conv = apply_parametrization_norm(nn.Conv1d(*args, **kwargs), norm)
116
+ self.norm = get_norm_module(self.conv, causal, norm, **norm_kwargs)
117
+ self.norm_type = norm
118
+
119
+ def forward(self, x):
120
+ x = self.conv(x)
121
+ x = self.norm(x)
122
+ return x
123
+
124
+
125
+ class NormConv2d(nn.Module):
126
+ """Wrapper around Conv2d and normalization applied to this conv
127
+ to provide a uniform interface across normalization approaches.
128
+ """
129
+ def __init__(self, *args, norm: str = 'none',
130
+ norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):
131
+ super().__init__()
132
+ self.conv = apply_parametrization_norm(nn.Conv2d(*args, **kwargs), norm)
133
+ self.norm = get_norm_module(self.conv, causal=False, norm=norm, **norm_kwargs)
134
+ self.norm_type = norm
135
+
136
+ def forward(self, x):
137
+ x = self.conv(x)
138
+ x = self.norm(x)
139
+ return x
140
+
141
+
142
+ class NormConvTranspose1d(nn.Module):
143
+ """Wrapper around ConvTranspose1d and normalization applied to this conv
144
+ to provide a uniform interface across normalization approaches.
145
+ """
146
+ def __init__(self, *args, causal: bool = False, norm: str = 'none',
147
+ norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):
148
+ super().__init__()
149
+ self.convtr = apply_parametrization_norm(nn.ConvTranspose1d(*args, **kwargs), norm)
150
+ self.norm = get_norm_module(self.convtr, causal, norm, **norm_kwargs)
151
+ self.norm_type = norm
152
+
153
+ def forward(self, x):
154
+ x = self.convtr(x)
155
+ x = self.norm(x)
156
+ return x
157
+
158
+
159
+ class NormConvTranspose2d(nn.Module):
160
+ """Wrapper around ConvTranspose2d and normalization applied to this conv
161
+ to provide a uniform interface across normalization approaches.
162
+ """
163
+ def __init__(self, *args, norm: str = 'none',
164
+ norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):
165
+ super().__init__()
166
+ self.convtr = apply_parametrization_norm(nn.ConvTranspose2d(*args, **kwargs), norm)
167
+ self.norm = get_norm_module(self.convtr, causal=False, norm=norm, **norm_kwargs)
168
+
169
+ def forward(self, x):
170
+ x = self.convtr(x)
171
+ x = self.norm(x)
172
+ return x
173
+
174
+
175
+ class SConv1d(nn.Module):
176
+ """Conv1d with some builtin handling of asymmetric or causal padding
177
+ and normalization.
178
+ """
179
+ def __init__(self, in_channels: int, out_channels: int,
180
+ kernel_size: int, stride: int = 1, dilation: int = 1,
181
+ groups: int = 1, bias: bool = True, causal: bool = False,
182
+ norm: str = 'none', norm_kwargs: tp.Dict[str, tp.Any] = {},
183
+ pad_mode: str = 'reflect'):
184
+ super().__init__()
185
+ # warn user on unusual setup between dilation and stride
186
+ if stride > 1 and dilation > 1:
187
+ warnings.warn('SConv1d has been initialized with stride > 1 and dilation > 1'
188
+ f' (kernel_size={kernel_size} stride={stride}, dilation={dilation}).')
189
+ self.conv = NormConv1d(in_channels, out_channels, kernel_size, stride,
190
+ dilation=dilation, groups=groups, bias=bias, causal=causal,
191
+ norm=norm, norm_kwargs=norm_kwargs)
192
+ self.causal = causal
193
+ self.pad_mode = pad_mode
194
+
195
+ def forward(self, x):
196
+ B, C, T = x.shape
197
+ kernel_size = self.conv.conv.kernel_size[0]
198
+ stride = self.conv.conv.stride[0]
199
+ dilation = self.conv.conv.dilation[0]
200
+ padding_total = (kernel_size - 1) * dilation - (stride - 1)
201
+ extra_padding = get_extra_padding_for_conv1d(x, kernel_size, stride, padding_total)
202
+ if self.causal:
203
+ # Left padding for causal
204
+ x = pad1d(x, (padding_total, extra_padding), mode=self.pad_mode)
205
+ else:
206
+ # Asymmetric padding required for odd strides
207
+ padding_right = padding_total // 2
208
+ padding_left = padding_total - padding_right
209
+ x = pad1d(x, (padding_left, padding_right + extra_padding), mode=self.pad_mode)
210
+ return self.conv(x)
211
+
212
+
213
+ class SConvTranspose1d(nn.Module):
214
+ """ConvTranspose1d with some builtin handling of asymmetric or causal padding
215
+ and normalization.
216
+ """
217
+ def __init__(self, in_channels: int, out_channels: int,
218
+ kernel_size: int, stride: int = 1, causal: bool = False,
219
+ norm: str = 'none', trim_right_ratio: float = 1.,
220
+ norm_kwargs: tp.Dict[str, tp.Any] = {}):
221
+ super().__init__()
222
+ self.convtr = NormConvTranspose1d(in_channels, out_channels, kernel_size, stride,
223
+ causal=causal, norm=norm, norm_kwargs=norm_kwargs)
224
+ self.causal = causal
225
+ self.trim_right_ratio = trim_right_ratio
226
+ assert self.causal or self.trim_right_ratio == 1., \
227
+ "`trim_right_ratio` != 1.0 only makes sense for causal convolutions"
228
+ assert self.trim_right_ratio >= 0. and self.trim_right_ratio <= 1.
229
+
230
+ def forward(self, x):
231
+ kernel_size = self.convtr.convtr.kernel_size[0]
232
+ stride = self.convtr.convtr.stride[0]
233
+ padding_total = kernel_size - stride
234
+
235
+ y = self.convtr(x)
236
+
237
+ # We will only trim fixed padding. Extra padding from `pad_for_conv1d` would be
238
+ # removed at the very end, when keeping only the right length for the output,
239
+ # as removing it here would require also passing the length at the matching layer
240
+ # in the encoder.
241
+ if self.causal:
242
+ # Trim the padding on the right according to the specified ratio
243
+ # if trim_right_ratio = 1.0, trim everything from right
244
+ padding_right = math.ceil(padding_total * self.trim_right_ratio)
245
+ padding_left = padding_total - padding_right
246
+ y = unpad1d(y, (padding_left, padding_right))
247
+ else:
248
+ # Asymmetric padding required for odd strides
249
+ padding_right = padding_total // 2
250
+ padding_left = padding_total - padding_right
251
+ y = unpad1d(y, (padding_left, padding_right))
252
+ return y
clean/audio/safeear/safeear/models/modules/lstm.py ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ """LSTM layers module."""
8
+
9
+ from torch import nn
10
+
11
+
12
+ class SLSTM(nn.Module):
13
+ """
14
+ LSTM without worrying about the hidden state, nor the layout of the data.
15
+ Expects input as convolutional layout.
16
+ """
17
+ def __init__(self, dimension: int, num_layers: int = 2, skip: bool = True, bidirectional: bool=False):
18
+ super().__init__()
19
+ self.bidirectional = bidirectional
20
+ self.skip = skip
21
+ self.lstm = nn.LSTM(dimension, dimension, num_layers, bidirectional=bidirectional)
22
+
23
+ def forward(self, x):
24
+ x = x.permute(2, 0, 1)
25
+ y, _ = self.lstm(x)
26
+ if self.bidirectional:
27
+ x = x.repeat(1, 1, 2)
28
+ if self.skip:
29
+ y = y + x
30
+ y = y.permute(1, 2, 0)
31
+ return y
32
+
clean/audio/safeear/safeear/models/modules/norm.py ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ """Normalization modules."""
8
+
9
+ import typing as tp
10
+
11
+ import einops
12
+ import torch
13
+ from torch import nn
14
+
15
+
16
+ class ConvLayerNorm(nn.LayerNorm):
17
+ """
18
+ Convolution-friendly LayerNorm that moves channels to last dimensions
19
+ before running the normalization and moves them back to original position right after.
20
+ """
21
+ def __init__(self, normalized_shape: tp.Union[int, tp.List[int], torch.Size], **kwargs):
22
+ super().__init__(normalized_shape, **kwargs)
23
+
24
+ def forward(self, x):
25
+ x = einops.rearrange(x, 'b ... t -> b t ...')
26
+ x = super().forward(x)
27
+ x = einops.rearrange(x, 'b t ... -> b ... t')
28
+ return
clean/audio/safeear/safeear/models/modules/quantization/__init__.py ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ # flake8: noqa
8
+ from .vq import QuantizedResult, ResidualVectorQuantizer
clean/audio/safeear/safeear/models/modules/quantization/ac.py ADDED
@@ -0,0 +1,292 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ """Arithmetic coder."""
8
+
9
+ import io
10
+ import math
11
+ import random
12
+ import typing as tp
13
+ import torch
14
+
15
+ from ..binary import BitPacker, BitUnpacker
16
+
17
+
18
+ def build_stable_quantized_cdf(pdf: torch.Tensor, total_range_bits: int,
19
+ roundoff: float = 1e-8, min_range: int = 2,
20
+ check: bool = True) -> torch.Tensor:
21
+ """Turn the given PDF into a quantized CDF that splits
22
+ [0, 2 ** self.total_range_bits - 1] into chunks of size roughly proportional
23
+ to the PDF.
24
+
25
+ Args:
26
+ pdf (torch.Tensor): probability distribution, shape should be `[N]`.
27
+ total_range_bits (int): see `ArithmeticCoder`, the typical range we expect
28
+ during the coding process is `[0, 2 ** total_range_bits - 1]`.
29
+ roundoff (float): will round the pdf up to that level to remove difference coming
30
+ from e.g. evaluating the Language Model on different architectures.
31
+ min_range (int): minimum range width. Should always be at least 2 for numerical
32
+ stability. Use this to avoid pathological behavior is a value
33
+ that is expected to be rare actually happens in real life.
34
+ check (bool): if True, checks that nothing bad happened, can be deactivated for speed.
35
+ """
36
+ pdf = pdf.detach()
37
+ if roundoff:
38
+ pdf = (pdf / roundoff).floor() * roundoff
39
+ # interpolate with uniform distribution to achieve desired minimum probability.
40
+ total_range = 2 ** total_range_bits
41
+ cardinality = len(pdf)
42
+ alpha = min_range * cardinality / total_range
43
+ assert alpha <= 1, "you must reduce min_range"
44
+ ranges = (((1 - alpha) * total_range) * pdf).floor().long()
45
+ ranges += min_range
46
+ quantized_cdf = torch.cumsum(ranges, dim=-1)
47
+ if min_range < 2:
48
+ raise ValueError("min_range must be at least 2.")
49
+ if check:
50
+ assert quantized_cdf[-1] <= 2 ** total_range_bits, quantized_cdf[-1]
51
+ if ((quantized_cdf[1:] - quantized_cdf[:-1]) < min_range).any() or quantized_cdf[0] < min_range:
52
+ raise ValueError("You must increase your total_range_bits.")
53
+ return quantized_cdf
54
+
55
+
56
+ class ArithmeticCoder:
57
+ """ArithmeticCoder,
58
+ Let us take a distribution `p` over `N` symbols, and assume we have a stream
59
+ of random variables `s_t` sampled from `p`. Let us assume that we have a budget
60
+ of `B` bits that we can afford to write on device. There are `2**B` possible numbers,
61
+ corresponding to the range `[0, 2 ** B - 1]`. We can map each of those number to a single
62
+ sequence `(s_t)` by doing the following:
63
+
64
+ 1) Initialize the current range to` [0 ** 2 B - 1]`.
65
+ 2) For each time step t, split the current range into contiguous chunks,
66
+ one for each possible outcome, with size roughly proportional to `p`.
67
+ For instance, if `p = [0.75, 0.25]`, and the range is `[0, 3]`, the chunks
68
+ would be `{[0, 2], [3, 3]}`.
69
+ 3) Select the chunk corresponding to `s_t`, and replace the current range with this.
70
+ 4) When done encoding all the values, just select any value remaining in the range.
71
+
72
+ You will notice that this procedure can fail: for instance if at any point in time
73
+ the range is smaller than `N`, then we can no longer assign a non-empty chunk to each
74
+ possible outcome. Intuitively, the more likely a value is, the less the range width
75
+ will reduce, and the longer we can go on encoding values. This makes sense: for any efficient
76
+ coding scheme, likely outcomes would take less bits, and more of them can be coded
77
+ with a fixed budget.
78
+
79
+ In practice, we do not know `B` ahead of time, but we have a way to inject new bits
80
+ when the current range decreases below a given limit (given by `total_range_bits`), without
81
+ having to redo all the computations. If we encode mostly likely values, we will seldom
82
+ need to inject new bits, but a single rare value can deplete our stock of entropy!
83
+
84
+ In this explanation, we assumed that the distribution `p` was constant. In fact, the present
85
+ code works for any sequence `(p_t)` possibly different for each timestep.
86
+ We also assume that `s_t ~ p_t`, but that doesn't need to be true, although the smaller
87
+ the KL between the true distribution and `p_t`, the most efficient the coding will be.
88
+
89
+ Args:
90
+ fo (IO[bytes]): file-like object to which the bytes will be written to.
91
+ total_range_bits (int): the range `M` described above is `2 ** total_range_bits.
92
+ Any time the current range width fall under this limit, new bits will
93
+ be injected to rescale the initial range.
94
+ """
95
+
96
+ def __init__(self, fo: tp.IO[bytes], total_range_bits: int = 24):
97
+ assert total_range_bits <= 30
98
+ self.total_range_bits = total_range_bits
99
+ self.packer = BitPacker(bits=1, fo=fo) # we push single bits at a time.
100
+ self.low: int = 0
101
+ self.high: int = 0
102
+ self.max_bit: int = -1
103
+ self._dbg: tp.List[tp.Any] = []
104
+ self._dbg2: tp.List[tp.Any] = []
105
+
106
+ @property
107
+ def delta(self) -> int:
108
+ """Return the current range width."""
109
+ return self.high - self.low + 1
110
+
111
+ def _flush_common_prefix(self):
112
+ # If self.low and self.high start with the sames bits,
113
+ # those won't change anymore as we always just increase the range
114
+ # by powers of 2, and we can flush them out to the bit stream.
115
+ assert self.high >= self.low, (self.low, self.high)
116
+ assert self.high < 2 ** (self.max_bit + 1)
117
+ while self.max_bit >= 0:
118
+ b1 = self.low >> self.max_bit
119
+ b2 = self.high >> self.max_bit
120
+ if b1 == b2:
121
+ self.low -= (b1 << self.max_bit)
122
+ self.high -= (b1 << self.max_bit)
123
+ assert self.high >= self.low, (self.high, self.low, self.max_bit)
124
+ assert self.low >= 0
125
+ self.max_bit -= 1
126
+ self.packer.push(b1)
127
+ else:
128
+ break
129
+
130
+ def push(self, symbol: int, quantized_cdf: torch.Tensor):
131
+ """Push the given symbol on the stream, flushing out bits
132
+ if possible.
133
+
134
+ Args:
135
+ symbol (int): symbol to encode with the AC.
136
+ quantized_cdf (torch.Tensor): use `build_stable_quantized_cdf`
137
+ to build this from your pdf estimate.
138
+ """
139
+ while self.delta < 2 ** self.total_range_bits:
140
+ self.low *= 2
141
+ self.high = self.high * 2 + 1
142
+ self.max_bit += 1
143
+
144
+ range_low = 0 if symbol == 0 else quantized_cdf[symbol - 1].item()
145
+ range_high = quantized_cdf[symbol].item() - 1
146
+ effective_low = int(math.ceil(range_low * (self.delta / (2 ** self.total_range_bits))))
147
+ effective_high = int(math.floor(range_high * (self.delta / (2 ** self.total_range_bits))))
148
+ assert self.low <= self.high
149
+ self.high = self.low + effective_high
150
+ self.low = self.low + effective_low
151
+ assert self.low <= self.high, (effective_low, effective_high, range_low, range_high)
152
+ self._dbg.append((self.low, self.high))
153
+ self._dbg2.append((self.low, self.high))
154
+ outs = self._flush_common_prefix()
155
+ assert self.low <= self.high
156
+ assert self.max_bit >= -1
157
+ assert self.max_bit <= 61, self.max_bit
158
+ return outs
159
+
160
+ def flush(self):
161
+ """Flush the remaining information to the stream.
162
+ """
163
+ while self.max_bit >= 0:
164
+ b1 = (self.low >> self.max_bit) & 1
165
+ self.packer.push(b1)
166
+ self.max_bit -= 1
167
+ self.packer.flush()
168
+
169
+
170
+ class ArithmeticDecoder:
171
+ """ArithmeticDecoder, see `ArithmeticCoder` for a detailed explanation.
172
+
173
+ Note that this must be called with **exactly** the same parameters and sequence
174
+ of quantized cdf as the arithmetic encoder or the wrong values will be decoded.
175
+
176
+ If the AC encoder current range is [L, H], with `L` and `H` having the some common
177
+ prefix (i.e. the same most significant bits), then this prefix will be flushed to the stream.
178
+ For instances, having read 3 bits `b1 b2 b3`, we know that `[L, H]` is contained inside
179
+ `[b1 b2 b3 0 ... 0 b1 b3 b3 1 ... 1]`. Now this specific sub-range can only be obtained
180
+ for a specific sequence of symbols and a binary-search allows us to decode those symbols.
181
+ At some point, the prefix `b1 b2 b3` will no longer be sufficient to decode new symbols,
182
+ and we will need to read new bits from the stream and repeat the process.
183
+
184
+ """
185
+ def __init__(self, fo: tp.IO[bytes], total_range_bits: int = 24):
186
+ self.total_range_bits = total_range_bits
187
+ self.low: int = 0
188
+ self.high: int = 0
189
+ self.current: int = 0
190
+ self.max_bit: int = -1
191
+ self.unpacker = BitUnpacker(bits=1, fo=fo) # we pull single bits at a time.
192
+ # Following is for debugging
193
+ self._dbg: tp.List[tp.Any] = []
194
+ self._dbg2: tp.List[tp.Any] = []
195
+ self._last: tp.Any = None
196
+
197
+ @property
198
+ def delta(self) -> int:
199
+ return self.high - self.low + 1
200
+
201
+ def _flush_common_prefix(self):
202
+ # Given the current range [L, H], if both have a common prefix,
203
+ # we know we can remove it from our representation to avoid handling large numbers.
204
+ while self.max_bit >= 0:
205
+ b1 = self.low >> self.max_bit
206
+ b2 = self.high >> self.max_bit
207
+ if b1 == b2:
208
+ self.low -= (b1 << self.max_bit)
209
+ self.high -= (b1 << self.max_bit)
210
+ self.current -= (b1 << self.max_bit)
211
+ assert self.high >= self.low
212
+ assert self.low >= 0
213
+ self.max_bit -= 1
214
+ else:
215
+ break
216
+
217
+ def pull(self, quantized_cdf: torch.Tensor) -> tp.Optional[int]:
218
+ """Pull a symbol, reading as many bits from the stream as required.
219
+ This returns `None` when the stream has been exhausted.
220
+
221
+ Args:
222
+ quantized_cdf (torch.Tensor): use `build_stable_quantized_cdf`
223
+ to build this from your pdf estimate. This must be **exatly**
224
+ the same cdf as the one used at encoding time.
225
+ """
226
+ while self.delta < 2 ** self.total_range_bits:
227
+ bit = self.unpacker.pull()
228
+ if bit is None:
229
+ return None
230
+ self.low *= 2
231
+ self.high = self.high * 2 + 1
232
+ self.current = self.current * 2 + bit
233
+ self.max_bit += 1
234
+
235
+ def bin_search(low_idx: int, high_idx: int):
236
+ # Binary search is not just for coding interviews :)
237
+ if high_idx < low_idx:
238
+ raise RuntimeError("Binary search failed")
239
+ mid = (low_idx + high_idx) // 2
240
+ range_low = quantized_cdf[mid - 1].item() if mid > 0 else 0
241
+ range_high = quantized_cdf[mid].item() - 1
242
+ effective_low = int(math.ceil(range_low * (self.delta / (2 ** self.total_range_bits))))
243
+ effective_high = int(math.floor(range_high * (self.delta / (2 ** self.total_range_bits))))
244
+ low = effective_low + self.low
245
+ high = effective_high + self.low
246
+ if self.current >= low:
247
+ if self.current <= high:
248
+ return (mid, low, high, self.current)
249
+ else:
250
+ return bin_search(mid + 1, high_idx)
251
+ else:
252
+ return bin_search(low_idx, mid - 1)
253
+
254
+ self._last = (self.low, self.high, self.current, self.max_bit)
255
+ sym, self.low, self.high, self.current = bin_search(0, len(quantized_cdf) - 1)
256
+ self._dbg.append((self.low, self.high, self.current))
257
+ self._flush_common_prefix()
258
+ self._dbg2.append((self.low, self.high, self.current))
259
+
260
+ return sym
261
+
262
+
263
+ def test():
264
+ torch.manual_seed(1234)
265
+ random.seed(1234)
266
+ for _ in range(4):
267
+ pdfs = []
268
+ cardinality = random.randrange(4000)
269
+ steps = random.randrange(100, 500)
270
+ fo = io.BytesIO()
271
+ encoder = ArithmeticCoder(fo)
272
+ symbols = []
273
+ for step in range(steps):
274
+ pdf = torch.softmax(torch.randn(cardinality), dim=0)
275
+ pdfs.append(pdf)
276
+ q_cdf = build_stable_quantized_cdf(pdf, encoder.total_range_bits)
277
+ symbol = torch.multinomial(pdf, 1).item()
278
+ symbols.append(symbol)
279
+ encoder.push(symbol, q_cdf)
280
+ encoder.flush()
281
+
282
+ fo.seek(0)
283
+ decoder = ArithmeticDecoder(fo)
284
+ for idx, (pdf, symbol) in enumerate(zip(pdfs, symbols)):
285
+ q_cdf = build_stable_quantized_cdf(pdf, encoder.total_range_bits)
286
+ decoded_symbol = decoder.pull(q_cdf)
287
+ assert decoded_symbol == symbol, idx
288
+ assert decoder.pull(torch.zeros(1)) is None
289
+
290
+
291
+ if __name__ == "__main__":
292
+ test()
clean/audio/safeear/safeear/models/modules/quantization/core_vq.py ADDED
@@ -0,0 +1,366 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+ #
7
+ # This implementation is inspired from
8
+ # https://github.com/lucidrains/vector-quantize-pytorch
9
+ # which is released under MIT License. Hereafter, the original license:
10
+ # MIT License
11
+ #
12
+ # Copyright (c) 2020 Phil Wang
13
+ #
14
+ # Permission is hereby granted, free of charge, to any person obtaining a copy
15
+ # of this software and associated documentation files (the "Software"), to deal
16
+ # in the Software without restriction, including without limitation the rights
17
+ # to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
18
+ # copies of the Software, and to permit persons to whom the Software is
19
+ # furnished to do so, subject to the following conditions:
20
+ #
21
+ # The above copyright notice and this permission notice shall be included in all
22
+ # copies or substantial portions of the Software.
23
+ #
24
+ # THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
25
+ # IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
26
+ # FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
27
+ # AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
28
+ # LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
29
+ # OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
30
+ # SOFTWARE.
31
+
32
+ """Core vector quantization implementation."""
33
+ import typing as tp
34
+
35
+ from einops import rearrange, repeat
36
+ import torch
37
+ from torch import nn
38
+ import torch.nn.functional as F
39
+
40
+ from .distrib import broadcast_tensors, rank
41
+
42
+
43
+ def default(val: tp.Any, d: tp.Any) -> tp.Any:
44
+ return val if val is not None else d
45
+
46
+
47
+ def ema_inplace(moving_avg, new, decay: float):
48
+ moving_avg.data.mul_(decay).add_(new, alpha=(1 - decay))
49
+
50
+
51
+ def laplace_smoothing(x, n_categories: int, epsilon: float = 1e-5):
52
+ return (x + epsilon) / (x.sum() + n_categories * epsilon)
53
+
54
+
55
+ def uniform_init(*shape: int):
56
+ t = torch.empty(shape)
57
+ nn.init.kaiming_uniform_(t)
58
+ return t
59
+
60
+
61
+ def sample_vectors(samples, num: int):
62
+ num_samples, device = samples.shape[0], samples.device
63
+
64
+ if num_samples >= num:
65
+ indices = torch.randperm(num_samples, device=device)[:num]
66
+ else:
67
+ indices = torch.randint(0, num_samples, (num,), device=device)
68
+
69
+ return samples[indices]
70
+
71
+
72
+ def kmeans(samples, num_clusters: int, num_iters: int = 10):
73
+ dim, dtype = samples.shape[-1], samples.dtype
74
+
75
+ means = sample_vectors(samples, num_clusters)
76
+
77
+ for _ in range(num_iters):
78
+ diffs = rearrange(samples, "n d -> n () d") - rearrange(
79
+ means, "c d -> () c d"
80
+ )
81
+ dists = -(diffs ** 2).sum(dim=-1)
82
+
83
+ buckets = dists.max(dim=-1).indices
84
+ bins = torch.bincount(buckets, minlength=num_clusters)
85
+ zero_mask = bins == 0
86
+ bins_min_clamped = bins.masked_fill(zero_mask, 1)
87
+
88
+ new_means = buckets.new_zeros(num_clusters, dim, dtype=dtype)
89
+ new_means.scatter_add_(0, repeat(buckets, "n -> n d", d=dim), samples)
90
+ new_means = new_means / bins_min_clamped[..., None]
91
+
92
+ means = torch.where(zero_mask[..., None], means, new_means)
93
+
94
+ return means, bins
95
+
96
+
97
+ class EuclideanCodebook(nn.Module):
98
+ """Codebook with Euclidean distance.
99
+ Args:
100
+ dim (int): Dimension.
101
+ codebook_size (int): Codebook size.
102
+ kmeans_init (bool): Whether to use k-means to initialize the codebooks.
103
+ If set to true, run the k-means algorithm on the first training batch and use
104
+ the learned centroids as initialization.
105
+ kmeans_iters (int): Number of iterations used for k-means algorithm at initialization.
106
+ decay (float): Decay for exponential moving average over the codebooks.
107
+ epsilon (float): Epsilon value for numerical stability.
108
+ threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes
109
+ that have an exponential moving average cluster size less than the specified threshold with
110
+ randomly selected vector from the current batch.
111
+ """
112
+ def __init__(
113
+ self,
114
+ dim: int,
115
+ codebook_size: int,
116
+ kmeans_init: int = False,
117
+ kmeans_iters: int = 10,
118
+ decay: float = 0.99,
119
+ epsilon: float = 1e-5,
120
+ threshold_ema_dead_code: int = 2,
121
+ ):
122
+ super().__init__()
123
+ self.decay = decay
124
+ init_fn: tp.Union[tp.Callable[..., torch.Tensor], tp.Any] = uniform_init if not kmeans_init else torch.zeros
125
+ embed = init_fn(codebook_size, dim)
126
+
127
+ self.codebook_size = codebook_size
128
+
129
+ self.kmeans_iters = kmeans_iters
130
+ self.epsilon = epsilon
131
+ self.threshold_ema_dead_code = threshold_ema_dead_code
132
+
133
+ self.register_buffer("inited", torch.Tensor([not kmeans_init]))
134
+ self.register_buffer("cluster_size", torch.zeros(codebook_size))
135
+ self.register_buffer("embed", embed)
136
+ self.register_buffer("embed_avg", embed.clone())
137
+
138
+ @torch.jit.ignore
139
+ def init_embed_(self, data):
140
+ if self.inited:
141
+ return
142
+
143
+ embed, cluster_size = kmeans(data, self.codebook_size, self.kmeans_iters)
144
+ self.embed.data.copy_(embed)
145
+ self.embed_avg.data.copy_(embed.clone())
146
+ self.cluster_size.data.copy_(cluster_size)
147
+ self.inited.data.copy_(torch.Tensor([True]))
148
+ # Make sure all buffers across workers are in sync after initialization
149
+ #broadcast_tensors(self.buffers())
150
+
151
+ def replace_(self, samples, mask):
152
+ modified_codebook = torch.where(
153
+ mask[..., None], sample_vectors(samples, self.codebook_size), self.embed
154
+ )
155
+ self.embed.data.copy_(modified_codebook)
156
+
157
+ def expire_codes_(self, batch_samples):
158
+ if self.threshold_ema_dead_code == 0:
159
+ return
160
+
161
+ expired_codes = self.cluster_size < self.threshold_ema_dead_code
162
+ if not torch.any(expired_codes):
163
+ return
164
+
165
+ batch_samples = rearrange(batch_samples, "... d -> (...) d")
166
+ self.replace_(batch_samples, mask=expired_codes)
167
+ #broadcast_tensors(self.buffers())
168
+
169
+ def preprocess(self, x):
170
+ x = rearrange(x, "... d -> (...) d")
171
+ return x
172
+
173
+ def quantize(self, x):
174
+ embed = self.embed.t()
175
+ dist = -(
176
+ x.pow(2).sum(1, keepdim=True)
177
+ - 2 * x @ embed
178
+ + embed.pow(2).sum(0, keepdim=True)
179
+ )
180
+ embed_ind = dist.max(dim=-1).indices
181
+ return embed_ind
182
+
183
+ def postprocess_emb(self, embed_ind, shape):
184
+ return embed_ind.view(*shape[:-1])
185
+
186
+ def dequantize(self, embed_ind):
187
+ quantize = F.embedding(embed_ind, self.embed)
188
+ return quantize
189
+
190
+ def encode(self, x):
191
+ shape = x.shape
192
+ # pre-process
193
+ x = self.preprocess(x)
194
+ # quantize
195
+ embed_ind = self.quantize(x)
196
+ # post-process
197
+ embed_ind = self.postprocess_emb(embed_ind, shape)
198
+ return embed_ind
199
+
200
+ def decode(self, embed_ind):
201
+ quantize = self.dequantize(embed_ind)
202
+ return quantize
203
+
204
+ def forward(self, x):
205
+ shape, dtype = x.shape, x.dtype
206
+ x = self.preprocess(x)
207
+
208
+ self.init_embed_(x)
209
+
210
+ embed_ind = self.quantize(x)
211
+ embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)
212
+ embed_ind = self.postprocess_emb(embed_ind, shape)
213
+ quantize = self.dequantize(embed_ind)
214
+
215
+ if self.training:
216
+ # We do the expiry of code at that point as buffers are in sync
217
+ # and all the workers will take the same decision.
218
+ self.expire_codes_(x)
219
+ ema_inplace(self.cluster_size, embed_onehot.sum(0), self.decay)
220
+ embed_sum = x.t() @ embed_onehot
221
+ ema_inplace(self.embed_avg, embed_sum.t(), self.decay)
222
+ cluster_size = (
223
+ laplace_smoothing(self.cluster_size, self.codebook_size, self.epsilon)
224
+ * self.cluster_size.sum()
225
+ )
226
+ embed_normalized = self.embed_avg / cluster_size.unsqueeze(1)
227
+ self.embed.data.copy_(embed_normalized)
228
+
229
+ return quantize, embed_ind
230
+
231
+
232
+ class VectorQuantization(nn.Module):
233
+ """Vector quantization implementation.
234
+ Currently supports only euclidean distance.
235
+ Args:
236
+ dim (int): Dimension
237
+ codebook_size (int): Codebook size
238
+ codebook_dim (int): Codebook dimension. If not defined, uses the specified dimension in dim.
239
+ decay (float): Decay for exponential moving average over the codebooks.
240
+ epsilon (float): Epsilon value for numerical stability.
241
+ kmeans_init (bool): Whether to use kmeans to initialize the codebooks.
242
+ kmeans_iters (int): Number of iterations used for kmeans initialization.
243
+ threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes
244
+ that have an exponential moving average cluster size less than the specified threshold with
245
+ randomly selected vector from the current batch.
246
+ commitment_weight (float): Weight for commitment loss.
247
+ """
248
+ def __init__(
249
+ self,
250
+ dim: int,
251
+ codebook_size: int,
252
+ codebook_dim: tp.Optional[int] = None,
253
+ decay: float = 0.99,
254
+ epsilon: float = 1e-5,
255
+ kmeans_init: bool = True,
256
+ kmeans_iters: int = 50,
257
+ threshold_ema_dead_code: int = 2,
258
+ commitment_weight: float = 1.,
259
+ ):
260
+ super().__init__()
261
+ _codebook_dim: int = default(codebook_dim, dim)
262
+
263
+ requires_projection = _codebook_dim != dim
264
+ self.project_in = (nn.Linear(dim, _codebook_dim) if requires_projection else nn.Identity())
265
+ self.project_out = (nn.Linear(_codebook_dim, dim) if requires_projection else nn.Identity())
266
+
267
+ self.epsilon = epsilon
268
+ self.commitment_weight = commitment_weight
269
+
270
+ self._codebook = EuclideanCodebook(dim=_codebook_dim, codebook_size=codebook_size,
271
+ kmeans_init=kmeans_init, kmeans_iters=kmeans_iters,
272
+ decay=decay, epsilon=epsilon,
273
+ threshold_ema_dead_code=threshold_ema_dead_code)
274
+ self.codebook_size = codebook_size
275
+
276
+ @property
277
+ def codebook(self):
278
+ return self._codebook.embed
279
+
280
+ def encode(self, x):
281
+ x = rearrange(x, "b d n -> b n d")
282
+ x = self.project_in(x)
283
+ embed_in = self._codebook.encode(x)
284
+ return embed_in
285
+
286
+ def decode(self, embed_ind):
287
+ quantize = self._codebook.decode(embed_ind)
288
+ quantize = self.project_out(quantize)
289
+ quantize = rearrange(quantize, "b n d -> b d n")
290
+ return quantize
291
+
292
+ def forward(self, x):
293
+ device = x.device
294
+ x = rearrange(x, "b d n -> b n d")
295
+ x = self.project_in(x)
296
+
297
+ quantize, embed_ind = self._codebook(x)
298
+
299
+ if self.training:
300
+ quantize = x + (quantize - x).detach()
301
+
302
+ loss = torch.tensor([0.0], device=device, requires_grad=self.training)
303
+
304
+ if self.training:
305
+ if self.commitment_weight > 0:
306
+ commit_loss = F.mse_loss(quantize.detach(), x)
307
+ loss = loss + commit_loss * self.commitment_weight
308
+
309
+ quantize = self.project_out(quantize)
310
+ quantize = rearrange(quantize, "b n d -> b d n")
311
+ return quantize, embed_ind, loss
312
+
313
+
314
+ class ResidualVectorQuantization(nn.Module):
315
+ """Residual vector quantization implementation.
316
+ Follows Algorithm 1. in https://arxiv.org/pdf/2107.03312.pdf
317
+ """
318
+ def __init__(self, *, num_quantizers, **kwargs):
319
+ super().__init__()
320
+ self.layers = nn.ModuleList(
321
+ [VectorQuantization(**kwargs) for _ in range(num_quantizers)]
322
+ )
323
+
324
+ def forward(self, x, n_q: tp.Optional[int] = None, layers: tp.Optional[list] = None):
325
+ quantized_out = 0.0
326
+ residual = x
327
+
328
+ all_losses = []
329
+ all_indices = []
330
+ out_quantized = []
331
+
332
+ n_q = n_q or len(self.layers)
333
+
334
+ for i, layer in enumerate(self.layers[:n_q]):
335
+ quantized, indices, loss = layer(residual)
336
+ residual = residual - quantized
337
+ quantized_out = quantized_out + quantized
338
+
339
+ all_indices.append(indices)
340
+ all_losses.append(loss)
341
+ if layers and i in layers:
342
+ out_quantized.append(quantized)
343
+
344
+ out_losses, out_indices = map(torch.stack, (all_losses, all_indices))
345
+ return quantized_out, out_indices, out_losses, out_quantized
346
+
347
+ def encode(self, x: torch.Tensor, n_q: tp.Optional[int] = None, st: tp.Optional[int]= None) -> torch.Tensor:
348
+ residual = x
349
+ all_indices = []
350
+ n_q = n_q or len(self.layers)
351
+ st = st or 0
352
+ for layer in self.layers[st:n_q]:
353
+ indices = layer.encode(residual)
354
+ quantized = layer.decode(indices)
355
+ residual = residual - quantized
356
+ all_indices.append(indices)
357
+ out_indices = torch.stack(all_indices)
358
+ return out_indices
359
+
360
+ def decode(self, q_indices: torch.Tensor, st: int=0) -> torch.Tensor:
361
+ quantized_out = torch.tensor(0.0, device=q_indices.device)
362
+ for i, indices in enumerate(q_indices):
363
+ layer = self.layers[st + i]
364
+ quantized = layer.decode(indices)
365
+ quantized_out = quantized_out + quantized
366
+ return quantized_out
clean/audio/safeear/safeear/models/modules/quantization/distrib.py ADDED
@@ -0,0 +1,126 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ """Torch distributed utilities."""
8
+
9
+ import typing as tp
10
+
11
+ import torch
12
+
13
+
14
+ def rank():
15
+ if torch.distributed.is_initialized():
16
+ return torch.distributed.get_rank()
17
+ else:
18
+ return 0
19
+
20
+
21
+ def world_size():
22
+ if torch.distributed.is_initialized():
23
+ return torch.distributed.get_world_size()
24
+ else:
25
+ return 1
26
+
27
+
28
+ def is_distributed():
29
+ return world_size() > 1
30
+
31
+
32
+ def all_reduce(tensor: torch.Tensor, op=torch.distributed.ReduceOp.SUM):
33
+ if is_distributed():
34
+ return torch.distributed.all_reduce(tensor, op)
35
+
36
+
37
+ def _is_complex_or_float(tensor):
38
+ return torch.is_floating_point(tensor) or torch.is_complex(tensor)
39
+
40
+
41
+ def _check_number_of_params(params: tp.List[torch.Tensor]):
42
+ # utility function to check that the number of params in all workers is the same,
43
+ # and thus avoid a deadlock with distributed all reduce.
44
+ if not is_distributed() or not params:
45
+ return
46
+ #print('params[0].device ', params[0].device)
47
+ tensor = torch.tensor([len(params)], device=params[0].device, dtype=torch.long)
48
+ all_reduce(tensor)
49
+ if tensor.item() != len(params) * world_size():
50
+ # If not all the workers have the same number, for at least one of them,
51
+ # this inequality will be verified.
52
+ raise RuntimeError(f"Mismatch in number of params: ours is {len(params)}, "
53
+ "at least one worker has a different one.")
54
+
55
+
56
+ def broadcast_tensors(tensors: tp.Iterable[torch.Tensor], src: int = 0):
57
+ """Broadcast the tensors from the given parameters to all workers.
58
+ This can be used to ensure that all workers have the same model to start with.
59
+ """
60
+ if not is_distributed():
61
+ return
62
+ tensors = [tensor for tensor in tensors if _is_complex_or_float(tensor)]
63
+ _check_number_of_params(tensors)
64
+ handles = []
65
+ for tensor in tensors:
66
+ # src = int(rank()) # added code
67
+ handle = torch.distributed.broadcast(tensor.data, src=src, async_op=True)
68
+ handles.append(handle)
69
+ for handle in handles:
70
+ handle.wait()
71
+
72
+
73
+ def sync_buffer(buffers, average=True):
74
+ """
75
+ Sync grad for buffers. If average is False, broadcast instead of averaging.
76
+ """
77
+ if not is_distributed():
78
+ return
79
+ handles = []
80
+ for buffer in buffers:
81
+ if torch.is_floating_point(buffer.data):
82
+ if average:
83
+ handle = torch.distributed.all_reduce(
84
+ buffer.data, op=torch.distributed.ReduceOp.SUM, async_op=True)
85
+ else:
86
+ handle = torch.distributed.broadcast(
87
+ buffer.data, src=0, async_op=True)
88
+ handles.append((buffer, handle))
89
+ for buffer, handle in handles:
90
+ handle.wait()
91
+ if average:
92
+ buffer.data /= world_size
93
+
94
+
95
+ def sync_grad(params):
96
+ """
97
+ Simpler alternative to DistributedDataParallel, that doesn't rely
98
+ on any black magic. For simple models it can also be as fast.
99
+ Just call this on your model parameters after the call to backward!
100
+ """
101
+ if not is_distributed():
102
+ return
103
+ handles = []
104
+ for p in params:
105
+ if p.grad is not None:
106
+ handle = torch.distributed.all_reduce(
107
+ p.grad.data, op=torch.distributed.ReduceOp.SUM, async_op=True)
108
+ handles.append((p, handle))
109
+ for p, handle in handles:
110
+ handle.wait()
111
+ p.grad.data /= world_size()
112
+
113
+
114
+ def average_metrics(metrics: tp.Dict[str, float], count=1.):
115
+ """Average a dictionary of metrics across all workers, using the optional
116
+ `count` as unormalized weight.
117
+ """
118
+ if not is_distributed():
119
+ return metrics
120
+ keys, values = zip(*metrics.items())
121
+ device = 'cuda' if torch.cuda.is_available() else 'cpu'
122
+ tensor = torch.tensor(list(values) + [1], device=device, dtype=torch.float32)
123
+ tensor *= count
124
+ all_reduce(tensor)
125
+ averaged = (tensor[:-1] / tensor[-1]).cpu().tolist()
126
+ return dict(zip(keys, averaged))
clean/audio/safeear/safeear/models/modules/quantization/vq.py ADDED
@@ -0,0 +1,108 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ """Residual vector quantizer implementation."""
8
+
9
+ from dataclasses import dataclass, field
10
+ import math
11
+ import typing as tp
12
+
13
+ import torch
14
+ from torch import nn
15
+
16
+ from .core_vq import ResidualVectorQuantization
17
+
18
+
19
+ @dataclass
20
+ class QuantizedResult:
21
+ quantized: torch.Tensor
22
+ codes: torch.Tensor
23
+ bandwidth: torch.Tensor # bandwidth in kb/s used, per batch item.
24
+ penalty: tp.Optional[torch.Tensor] = None
25
+ metrics: dict = field(default_factory=dict)
26
+
27
+
28
+ class ResidualVectorQuantizer(nn.Module):
29
+ """Residual Vector Quantizer.
30
+ Args:
31
+ dimension (int): Dimension of the codebooks.
32
+ n_q (int): Number of residual vector quantizers used.
33
+ bins (int): Codebook size.
34
+ decay (float): Decay for exponential moving average over the codebooks.
35
+ kmeans_init (bool): Whether to use kmeans to initialize the codebooks.
36
+ kmeans_iters (int): Number of iterations used for kmeans initialization.
37
+ threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes
38
+ that have an exponential moving average cluster size less than the specified threshold with
39
+ randomly selected vector from the current batch.
40
+ """
41
+ def __init__(
42
+ self,
43
+ dimension: int = 256,
44
+ n_q: int = 8,
45
+ bins: int = 1024,
46
+ decay: float = 0.99,
47
+ kmeans_init: bool = True,
48
+ kmeans_iters: int = 50,
49
+ threshold_ema_dead_code: int = 2,
50
+ ):
51
+ super().__init__()
52
+ self.n_q = n_q
53
+ self.dimension = dimension
54
+ self.bins = bins
55
+ self.decay = decay
56
+ self.kmeans_init = kmeans_init
57
+ self.kmeans_iters = kmeans_iters
58
+ self.threshold_ema_dead_code = threshold_ema_dead_code
59
+ self.vq = ResidualVectorQuantization(
60
+ dim=self.dimension,
61
+ codebook_size=self.bins,
62
+ num_quantizers=self.n_q,
63
+ decay=self.decay,
64
+ kmeans_init=self.kmeans_init,
65
+ kmeans_iters=self.kmeans_iters,
66
+ threshold_ema_dead_code=self.threshold_ema_dead_code,
67
+ )
68
+
69
+ def forward(self, x: torch.Tensor, n_q: tp.Optional[int] = None, layers: tp.Optional[list] = None) -> QuantizedResult:
70
+ """Residual vector quantization on the given input tensor.
71
+ Args:
72
+ x (torch.Tensor): Input tensor.
73
+ n_q (int): Number of quantizer used to quantize. Default: All quantizers.
74
+ layers (list): Layer that need to return quantized. Defalt: None.
75
+ Returns:
76
+ QuantizedResult:
77
+ The quantized (or approximately quantized) representation with
78
+ the associated numbert quantizers and layer quantized required to return.
79
+ """
80
+ n_q = n_q if n_q else self.n_q
81
+ if layers and max(layers) >= n_q:
82
+ raise ValueError(f'Last layer index in layers: A {max(layers)}. Number of quantizers in RVQ: B {self.n_q}. A must less than B.')
83
+ quantized, codes, commit_loss, quantized_list = self.vq(x, n_q=n_q, layers=layers)
84
+ return quantized, codes, torch.mean(commit_loss), quantized_list
85
+
86
+
87
+ def encode(self, x: torch.Tensor, n_q: tp.Optional[int] = None, st: tp.Optional[int] = None) -> torch.Tensor:
88
+ """Encode a given input tensor with the specified sample rate at the given bandwidth.
89
+ The RVQ encode method sets the appropriate number of quantizer to use
90
+ and returns indices for each quantizer.
91
+ Args:
92
+ x (torch.Tensor): Input tensor.
93
+ n_q (int): Number of quantizer used to quantize. Default: All quantizers.
94
+ st (int): Start to encode input from which layers. Default: 0.
95
+ """
96
+ n_q = n_q if n_q else self.n_q
97
+ st = st or 0
98
+ codes = self.vq.encode(x, n_q=n_q, st=st)
99
+ return codes
100
+
101
+ def decode(self, codes: torch.Tensor, st: int = 0) -> torch.Tensor:
102
+ """Decode the given codes to the quantized representation.
103
+ Args:
104
+ codes (torch.Tensor): Input indices for each quantizer.
105
+ st (int): Start to decode input codes from which layers. Default: 0.
106
+ """
107
+ quantized = self.vq.decode(codes, st=st)
108
+ return quantized
clean/audio/safeear/safeear/models/modules/seanet.py ADDED
@@ -0,0 +1,275 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+ # All rights reserved.
3
+ #
4
+ # This source code is licensed under the license found in the
5
+ # LICENSE file in the root directory of this source tree.
6
+
7
+ """Encodec SEANet-based encoder and decoder implementation."""
8
+
9
+ import typing as tp
10
+
11
+ import numpy as np
12
+ import torch.nn as nn
13
+ import torch
14
+
15
+ from . import (
16
+ SConv1d,
17
+ SConvTranspose1d,
18
+ SLSTM
19
+ )
20
+
21
+
22
+ @torch.jit.script
23
+ def snake(x, alpha):
24
+ shape = x.shape
25
+ x = x.reshape(shape[0], shape[1], -1)
26
+ x = x + (alpha + 1e-9).reciprocal() * torch.sin(alpha * x).pow(2)
27
+ x = x.reshape(shape)
28
+ return x
29
+
30
+
31
+ class Snake1d(nn.Module):
32
+ def __init__(self, channels):
33
+ super().__init__()
34
+ self.alpha = nn.Parameter(torch.ones(1, channels, 1))
35
+
36
+ def forward(self, x):
37
+ return snake(x, self.alpha)
38
+
39
+ class SEANetResnetBlock(nn.Module):
40
+ """Residual block from SEANet model.
41
+ Args:
42
+ dim (int): Dimension of the input/output
43
+ kernel_sizes (list): List of kernel sizes for the convolutions.
44
+ dilations (list): List of dilations for the convolutions.
45
+ activation (str): Activation function.
46
+ activation_params (dict): Parameters to provide to the activation function
47
+ norm (str): Normalization method.
48
+ norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution.
49
+ causal (bool): Whether to use fully causal convolution.
50
+ pad_mode (str): Padding mode for the convolutions.
51
+ compress (int): Reduced dimensionality in residual branches (from Demucs v3)
52
+ true_skip (bool): Whether to use true skip connection or a simple convolution as the skip connection.
53
+ """
54
+ def __init__(self, dim: int, kernel_sizes: tp.List[int] = [3, 1], dilations: tp.List[int] = [1, 1],
55
+ activation: str = 'ELU', activation_params: dict = {'alpha': 1.0},
56
+ norm: str = 'weight_norm', norm_params: tp.Dict[str, tp.Any] = {}, causal: bool = False,
57
+ pad_mode: str = 'reflect', compress: int = 2, true_skip: bool = True):
58
+ super().__init__()
59
+ assert len(kernel_sizes) == len(dilations), 'Number of kernel sizes should match number of dilations'
60
+ act = getattr(nn, activation) if activation != 'Snake' else Snake1d
61
+ hidden = dim // compress
62
+ block = []
63
+ for i, (kernel_size, dilation) in enumerate(zip(kernel_sizes, dilations)):
64
+ in_chs = dim if i == 0 else hidden
65
+ out_chs = dim if i == len(kernel_sizes) - 1 else hidden
66
+ block += [
67
+ act(**activation_params) if activation != 'Snake' else act(in_chs),
68
+ SConv1d(in_chs, out_chs, kernel_size=kernel_size, dilation=dilation,
69
+ norm=norm, norm_kwargs=norm_params,
70
+ causal=causal, pad_mode=pad_mode),
71
+ ]
72
+ self.block = nn.Sequential(*block)
73
+ self.shortcut: nn.Module
74
+ if true_skip:
75
+ self.shortcut = nn.Identity()
76
+ else:
77
+ self.shortcut = SConv1d(dim, dim, kernel_size=1, norm=norm, norm_kwargs=norm_params,
78
+ causal=causal, pad_mode=pad_mode)
79
+
80
+ def forward(self, x):
81
+ return self.shortcut(x) + self.block(x)
82
+
83
+
84
+
85
+ class SEANetEncoder(nn.Module):
86
+ """SEANet encoder.
87
+ Args:
88
+ channels (int): Audio channels.
89
+ dimension (int): Intermediate representation dimension.
90
+ n_filters (int): Base width for the model.
91
+ n_residual_layers (int): nb of residual layers.
92
+ ratios (Sequence[int]): kernel size and stride ratios. The encoder uses downsampling ratios instead of
93
+ upsampling ratios, hence it will use the ratios in the reverse order to the ones specified here
94
+ that must match the decoder order
95
+ activation (str): Activation function.
96
+ activation_params (dict): Parameters to provide to the activation function
97
+ norm (str): Normalization method.
98
+ norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution.
99
+ kernel_size (int): Kernel size for the initial convolution.
100
+ last_kernel_size (int): Kernel size for the initial convolution.
101
+ residual_kernel_size (int): Kernel size for the residual layers.
102
+ dilation_base (int): How much to increase the dilation with each layer.
103
+ causal (bool): Whether to use fully causal convolution.
104
+ pad_mode (str): Padding mode for the convolutions.
105
+ true_skip (bool): Whether to use true skip connection or a simple
106
+ (streamable) convolution as the skip connection in the residual network blocks.
107
+ compress (int): Reduced dimensionality in residual branches (from Demucs v3).
108
+ lstm (int): Number of LSTM layers at the end of the encoder.
109
+ """
110
+ def __init__(self, channels: int = 1, dimension: int = 128, n_filters: int = 32, n_residual_layers: int = 1,
111
+ ratios: tp.List[int] = [8, 5, 4, 2], activation: str = 'ELU', activation_params: dict = {'alpha': 1.0},
112
+ norm: str = 'weight_norm', norm_params: tp.Dict[str, tp.Any] = {}, kernel_size: int = 7,
113
+ last_kernel_size: int = 7, residual_kernel_size: int = 3, dilation_base: int = 2, causal: bool = False,
114
+ pad_mode: str = 'reflect', true_skip: bool = False, compress: int = 2, lstm: int = 2, bidirectional:bool = False):
115
+ super().__init__()
116
+ self.channels = channels
117
+ self.dimension = dimension
118
+ self.n_filters = n_filters
119
+ self.ratios = list(reversed(ratios))
120
+ del ratios
121
+ self.n_residual_layers = n_residual_layers
122
+ self.hop_length = np.prod(self.ratios) # 计算乘积
123
+
124
+ act = getattr(nn, activation) if activation != 'Snake' else Snake1d
125
+ mult = 1
126
+ model: tp.List[nn.Module] = [
127
+ SConv1d(channels, mult * n_filters, kernel_size, norm=norm, norm_kwargs=norm_params,
128
+ causal=causal, pad_mode=pad_mode)
129
+ ]
130
+ # Downsample to raw audio scale
131
+ for i, ratio in enumerate(self.ratios):
132
+ # Add residual layers
133
+ for j in range(n_residual_layers):
134
+ model += [
135
+ SEANetResnetBlock(mult * n_filters, kernel_sizes=[residual_kernel_size, 1],
136
+ dilations=[dilation_base ** j, 1],
137
+ norm=norm, norm_params=norm_params,
138
+ activation=activation, activation_params=activation_params,
139
+ causal=causal, pad_mode=pad_mode, compress=compress, true_skip=true_skip)]
140
+
141
+ # Add downsampling layers
142
+ model += [
143
+ act(**activation_params) if activation != 'Snake' else act(mult * n_filters),
144
+ SConv1d(mult * n_filters, mult * n_filters * 2,
145
+ kernel_size=ratio * 2, stride=ratio,
146
+ norm=norm, norm_kwargs=norm_params,
147
+ causal=causal, pad_mode=pad_mode),
148
+ ]
149
+ mult *= 2
150
+
151
+ if lstm:
152
+ model += [SLSTM(mult * n_filters, num_layers=lstm, bidirectional=bidirectional)]
153
+
154
+ mult = mult * 2 if bidirectional else mult
155
+ model += [
156
+ act(**activation_params) if activation != 'Snake' else act(mult * n_filters),
157
+ SConv1d(mult * n_filters, dimension, last_kernel_size, norm=norm, norm_kwargs=norm_params,
158
+ causal=causal, pad_mode=pad_mode)
159
+ ]
160
+
161
+ self.model = nn.Sequential(*model)
162
+
163
+ def forward(self, x):
164
+ return self.model(x)
165
+
166
+
167
+ class SEANetDecoder(nn.Module):
168
+ """SEANet decoder.
169
+ Args:
170
+ channels (int): Audio channels.
171
+ dimension (int): Intermediate representation dimension.
172
+ n_filters (int): Base width for the model.
173
+ n_residual_layers (int): nb of residual layers.
174
+ ratios (Sequence[int]): kernel size and stride ratios
175
+ activation (str): Activation function.
176
+ activation_params (dict): Parameters to provide to the activation function
177
+ final_activation (str): Final activation function after all convolutions.
178
+ final_activation_params (dict): Parameters to provide to the activation function
179
+ norm (str): Normalization method.
180
+ norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution.
181
+ kernel_size (int): Kernel size for the initial convolution.
182
+ last_kernel_size (int): Kernel size for the initial convolution.
183
+ residual_kernel_size (int): Kernel size for the residual layers.
184
+ dilation_base (int): How much to increase the dilation with each layer.
185
+ causal (bool): Whether to use fully causal convolution.
186
+ pad_mode (str): Padding mode for the convolutions.
187
+ true_skip (bool): Whether to use true skip connection or a simple
188
+ (streamable) convolution as the skip connection in the residual network blocks.
189
+ compress (int): Reduced dimensionality in residual branches (from Demucs v3).
190
+ lstm (int): Number of LSTM layers at the end of the encoder.
191
+ trim_right_ratio (float): Ratio for trimming at the right of the transposed convolution under the causal setup.
192
+ If equal to 1.0, it means that all the trimming is done at the right.
193
+ """
194
+ def __init__(self, channels: int = 1, dimension: int = 128, n_filters: int = 32, n_residual_layers: int = 1,
195
+ ratios: tp.List[int] = [8, 5, 4, 2], activation: str = 'ELU', activation_params: dict = {'alpha': 1.0},
196
+ final_activation: tp.Optional[str] = None, final_activation_params: tp.Optional[dict] = None,
197
+ norm: str = 'weight_norm', norm_params: tp.Dict[str, tp.Any] = {}, kernel_size: int = 7,
198
+ last_kernel_size: int = 7, residual_kernel_size: int = 3, dilation_base: int = 2, causal: bool = False,
199
+ pad_mode: str = 'reflect', true_skip: bool = False, compress: int = 2, lstm: int = 2,
200
+ trim_right_ratio: float = 1.0, bidirectional:bool = False):
201
+ super().__init__()
202
+ self.dimension = dimension
203
+ self.channels = channels
204
+ self.n_filters = n_filters
205
+ self.ratios = ratios
206
+ del ratios
207
+ self.n_residual_layers = n_residual_layers
208
+ self.hop_length = np.prod(self.ratios)
209
+
210
+ act = getattr(nn, activation) if activation != 'Snake' else Snake1d
211
+ mult = int(2 ** len(self.ratios))
212
+ model: tp.List[nn.Module] = [
213
+ SConv1d(dimension, mult * n_filters, kernel_size, norm=norm, norm_kwargs=norm_params,
214
+ causal=causal, pad_mode=pad_mode)
215
+ ]
216
+
217
+ if lstm:
218
+ model += [SLSTM(mult * n_filters, num_layers=lstm, bidirectional=bidirectional)]
219
+
220
+ # Upsample to raw audio scale
221
+ for i, ratio in enumerate(self.ratios):
222
+ # Add upsampling layers
223
+ model += [
224
+ act(**activation_params) if activation != 'Snake' else act(mult * n_filters),
225
+ SConvTranspose1d(mult * n_filters, mult * n_filters // 2,
226
+ kernel_size=ratio * 2, stride=ratio,
227
+ norm=norm, norm_kwargs=norm_params,
228
+ causal=causal, trim_right_ratio=trim_right_ratio),
229
+ ]
230
+ # Add residual layers
231
+ for j in range(n_residual_layers):
232
+ model += [
233
+ SEANetResnetBlock(mult * n_filters // 2, kernel_sizes=[residual_kernel_size, 1],
234
+ dilations=[dilation_base ** j, 1],
235
+ activation=activation, activation_params=activation_params,
236
+ norm=norm, norm_params=norm_params, causal=causal,
237
+ pad_mode=pad_mode, compress=compress, true_skip=true_skip)]
238
+
239
+ mult //= 2
240
+
241
+ # Add final layers
242
+ model += [
243
+ act(**activation_params) if activation != 'Snake' else act(n_filters),
244
+ SConv1d(n_filters, channels, last_kernel_size, norm=norm, norm_kwargs=norm_params,
245
+ causal=causal, pad_mode=pad_mode)
246
+ ]
247
+ # Add optional final activation to decoder (eg. tanh)
248
+ if final_activation is not None:
249
+ final_act = getattr(nn, final_activation)
250
+ final_activation_params = final_activation_params or {}
251
+ model += [
252
+ final_act(**final_activation_params)
253
+ ]
254
+ self.model = nn.Sequential(*model)
255
+
256
+ def forward(self, z):
257
+ y = self.model(z)
258
+ return y
259
+
260
+
261
+ def test():
262
+ import torch
263
+ encoder = SEANetEncoder()
264
+ decoder = SEANetDecoder()
265
+ x = torch.randn(1, 1, 24000)
266
+ z = encoder(x)
267
+ print('z ', z.shape)
268
+ assert 1==2
269
+ assert list(z.shape) == [1, 128, 75], z.shape
270
+ y = decoder(z)
271
+ assert y.shape == x.shape, (x.shape, y.shape)
272
+
273
+
274
+ if __name__ == '__main__':
275
+ test()
clean/audio/safeear/safeear/models/safeear.py ADDED
@@ -0,0 +1,959 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+ import torch.nn as nn
4
+ import torch.nn.functional as F
5
+ from torch.nn import Module, ModuleList, Linear, Dropout, LayerNorm, Identity, Parameter, init
6
+ from timm.models.layers import trunc_normal_, DropPath
7
+ import random
8
+ from typing import Union
9
+ import numpy as np
10
+ import math
11
+ from torch import Tensor
12
+
13
+ def conv3x3(in_planes, out_planes, stride=1):
14
+ """3x3 convolution with padding"""
15
+ return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,
16
+ padding=1, bias=False)
17
+
18
+ class SELayer(nn.Module):
19
+ def __init__(self, channel, reduction=16):
20
+ super(SELayer, self).__init__()
21
+ # print('se reduction: ', reduction)
22
+ # print(channel // reduction)
23
+ self.avg_pool = nn.AdaptiveAvgPool2d(1) # F_squeeze
24
+ self.fc = nn.Sequential(
25
+ nn.Linear(channel, channel // reduction, bias=False),
26
+ nn.ReLU(inplace=True),
27
+ nn.Linear(channel // reduction, channel, bias=False),
28
+ nn.Sigmoid()
29
+ )
30
+
31
+ def forward(self, x): # x: B*C*D*T
32
+ b, c, _, _ = x.size()
33
+ y = self.avg_pool(x).view(b, c)
34
+ y = self.fc(y).view(b, c, 1, 1)
35
+ return x * y.expand_as(x)
36
+
37
+ class BasicBlock(nn.Module):
38
+ expansion = 1
39
+
40
+ def __init__(self, inplanes, planes, stride=1, downsample=None):
41
+ super(BasicBlock, self).__init__()
42
+ self.conv1 = conv3x3(inplanes, planes, stride)
43
+ self.bn1 = nn.BatchNorm2d(planes)
44
+ self.relu = nn.ReLU(inplace=True)
45
+ self.conv2 = conv3x3(planes, planes)
46
+ self.bn2 = nn.BatchNorm2d(planes)
47
+ self.downsample = downsample
48
+ self.stride = stride
49
+
50
+ def forward(self, x):
51
+ residual = x
52
+
53
+ out = self.conv1(x)
54
+ out = self.bn1(out)
55
+ out = self.relu(out)
56
+
57
+ out = self.conv2(out)
58
+ out = self.bn2(out)
59
+
60
+ if self.downsample is not None:
61
+ residual = self.downsample(x)
62
+
63
+ out += residual
64
+ out = self.relu(out)
65
+
66
+ return out
67
+
68
+ class SEBasicBlock(nn.Module):
69
+ expansion = 1
70
+
71
+ def __init__(self, inplanes, planes, stride=1, downsample=None, reduction=16):
72
+ super(SEBasicBlock, self).__init__()
73
+ self.conv1 = conv3x3(inplanes, planes, stride)
74
+ self.bn1 = nn.BatchNorm2d(planes)
75
+ self.relu = nn.ReLU(inplace=True)
76
+ self.conv2 = conv3x3(planes, planes, 1)
77
+ self.bn2 = nn.BatchNorm2d(planes)
78
+ self.se = SELayer(planes, reduction)
79
+ self.downsample = downsample
80
+ self.stride = stride
81
+
82
+ def forward(self, x):
83
+ residual = x
84
+ out = self.conv1(x)
85
+ out = self.bn1(out)
86
+ out = self.relu(out)
87
+
88
+ out = self.conv2(out)
89
+ out = self.bn2(out)
90
+ out = self.se(out)
91
+
92
+ if self.downsample is not None:
93
+ residual = self.downsample(x)
94
+
95
+ out += residual
96
+ out = self.relu(out)
97
+
98
+ return out
99
+
100
+ class Bottleneck(nn.Module):
101
+ expansion = 2
102
+
103
+ def __init__(self, inplanes, planes, stride=1, downsample=None):
104
+ super(Bottleneck, self).__init__()
105
+ self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=1, bias=False)
106
+ self.bn1 = nn.BatchNorm2d(planes)
107
+ self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=stride, padding=1, bias=False)
108
+ self.bn2 = nn.BatchNorm2d(planes)
109
+ self.conv3 = nn.Conv2d(planes, planes * self.expansion, kernel_size=1, bias=False)
110
+ self.bn3 = nn.BatchNorm2d(planes * self.expansion)
111
+ self.relu = nn.ReLU(inplace=True)
112
+ self.downsample = downsample
113
+ self.stride = stride
114
+
115
+ def forward(self, x):
116
+ residual = x
117
+
118
+ out = self.conv1(x)
119
+ out = self.bn1(out)
120
+ out = self.relu(out)
121
+
122
+ out = self.conv2(out)
123
+ out = self.bn2(out)
124
+ out = self.relu(out)
125
+
126
+ out = self.conv3(out)
127
+ out = self.bn3(out)
128
+
129
+ if self.downsample is not None:
130
+ residual = self.downsample(x)
131
+
132
+ out += residual
133
+ out = self.relu(out)
134
+
135
+ return out
136
+
137
+ class SEBottleneck(nn.Module):
138
+ expansion = 2
139
+
140
+ def __init__(self, inplanes, planes, stride=1, downsample=None, reduction=16):
141
+ super(SEBottleneck, self).__init__()
142
+ self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=1, bias=False)
143
+ self.bn1 = nn.BatchNorm2d(planes)
144
+ self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=stride, padding=1, bias=False)
145
+ self.bn2 = nn.BatchNorm2d(planes)
146
+ self.conv3 = nn.Conv2d(planes, planes * self.expansion, kernel_size=1, bias=False)
147
+ self.bn3 = nn.BatchNorm2d(planes * self.expansion)
148
+ self.relu = nn.ReLU(inplace=True)
149
+ self.se = SELayer(planes * self.expansion, reduction)
150
+ self.downsample = downsample
151
+ self.stride = stride
152
+
153
+ def forward(self, x):
154
+ residual = x
155
+
156
+ out = self.conv1(x)
157
+ out = self.bn1(out)
158
+ out = self.relu(out)
159
+
160
+ out = self.conv2(out)
161
+ out = self.bn2(out)
162
+ out = self.relu(out)
163
+
164
+ out = self.conv3(out)
165
+ out = self.bn3(out)
166
+ out = self.se(out)
167
+
168
+ if self.downsample is not None:
169
+ residual = self.downsample(x)
170
+
171
+ out += residual
172
+ out = self.relu(out)
173
+
174
+ return out
175
+
176
+
177
+ class Bottle2neck(nn.Module):
178
+ expansion = 2
179
+
180
+ def __init__(self,
181
+ inplanes,
182
+ planes,
183
+ stride=1,
184
+ downsample=None,
185
+ baseWidth=26,
186
+ scale=4,
187
+ stype='normal'):
188
+ """ Constructor
189
+ Args:
190
+ inplanes: input channel dimensionality
191
+ planes: output channel dimensionality
192
+ stride: conv stride. Replaces pooling layer.
193
+ downsample: None when stride = 1
194
+ baseWidth: basic width of conv3x3
195
+ scale: number of scale.
196
+ type: 'normal': normal set. 'stage': first block of a new stage.
197
+ """
198
+ super(Bottle2neck, self).__init__()
199
+
200
+ width = int(math.floor(planes * (baseWidth / 64.0)))
201
+ self.conv1 = nn.Conv2d(inplanes,
202
+ width * scale,
203
+ kernel_size=1,
204
+ bias=False)
205
+ self.bn1 = nn.BatchNorm2d(width * scale)
206
+
207
+ if scale == 1:
208
+ self.nums = 1
209
+ else:
210
+ self.nums = scale - 1
211
+ if stype == 'stage':
212
+ self.pool = nn.AvgPool2d(kernel_size=3, stride=stride, padding=1)
213
+ convs = []
214
+ bns = []
215
+ for i in range(self.nums):
216
+ convs.append(
217
+ nn.Conv2d(width,
218
+ width,
219
+ kernel_size=3,
220
+ stride=stride,
221
+ padding=1,
222
+ bias=False))
223
+ bns.append(nn.BatchNorm2d(width))
224
+ self.convs = nn.ModuleList(convs)
225
+ self.bns = nn.ModuleList(bns)
226
+
227
+ self.conv3 = nn.Conv2d(width * scale,
228
+ planes * self.expansion,
229
+ kernel_size=1,
230
+ bias=False)
231
+ self.bn3 = nn.BatchNorm2d(planes * self.expansion)
232
+
233
+ self.relu = nn.ReLU(inplace=True)
234
+ if stride != 1 or inplanes != planes * self.expansion:
235
+ downsample = nn.Sequential(
236
+ nn.AvgPool2d(kernel_size=stride,
237
+ stride=stride,
238
+ ceil_mode=True,
239
+ count_include_pad=False),
240
+ nn.Conv2d(inplanes,
241
+ planes * self.expansion,
242
+ kernel_size=1,
243
+ stride=1,
244
+ bias=False),
245
+ nn.BatchNorm2d(planes * self.expansion),
246
+ )
247
+ self.downsample = downsample
248
+ self.stype = stype
249
+ self.scale = scale
250
+ self.width = width
251
+
252
+ def forward(self, x):
253
+ residual = x
254
+
255
+ out = self.conv1(x)
256
+ out = self.bn1(out)
257
+ out = self.relu(out)
258
+
259
+ spx = torch.split(out, self.width, 1)
260
+ for i in range(self.nums):
261
+ if i == 0 or self.stype == 'stage':
262
+ sp = spx[i]
263
+ else:
264
+ sp = sp + spx[i]
265
+ sp = self.convs[i](sp)
266
+ sp = self.relu(self.bns[i](sp))
267
+ if i == 0:
268
+ out = sp
269
+ else:
270
+ out = torch.cat((out, sp), 1)
271
+ if self.scale != 1 and self.stype == 'normal':
272
+ out = torch.cat((out, spx[self.nums]), 1)
273
+ elif self.scale != 1 and self.stype == 'stage':
274
+ out = torch.cat((out, self.pool(spx[self.nums])), 1)
275
+
276
+ out = self.conv3(out)
277
+ out = self.bn3(out)
278
+
279
+ if self.downsample is not None:
280
+ residual = self.downsample(x)
281
+
282
+ out += residual
283
+ out = self.relu(out)
284
+
285
+ return out
286
+
287
+ class SEBottle2neck(nn.Module):
288
+ expansion = 1
289
+
290
+ def __init__(self,
291
+ inplanes,
292
+ planes,
293
+ stride=1,
294
+ kernel_size = 1,
295
+ padding=1,
296
+ downsample=None,
297
+ baseWidth=26,
298
+ scale=4,
299
+ stype='normal'):
300
+ """ Constructor
301
+ Args:
302
+ inplanes: input channel dimensionality
303
+ planes: output channel dimensionality
304
+ stride: conv stride. Replaces pooling layer.
305
+ downsample: None when stride = 1
306
+ baseWidth: basic width of conv3x3
307
+ scale: number of scale.
308
+ type: 'normal': normal set. 'stage': first block of a new stage.
309
+ """
310
+ super(SEBottle2neck, self).__init__()
311
+
312
+ width = int(math.floor(planes * (baseWidth / 64.0)))
313
+ self.conv1 = nn.Conv2d(inplanes,
314
+ width * scale,
315
+ kernel_size=kernel_size,
316
+ padding=padding,
317
+ bias=False)
318
+ self.bn1 = nn.BatchNorm2d(width * scale)
319
+
320
+ if scale == 1:
321
+ self.nums = 1
322
+ else:
323
+ self.nums = scale - 1
324
+ if stype == 'stage':
325
+ self.pool = nn.AvgPool2d(kernel_size=3, stride=stride, padding=1)
326
+ convs = []
327
+ bns = []
328
+ for i in range(self.nums):
329
+ convs.append(
330
+ nn.Conv2d(width,
331
+ width,
332
+ kernel_size=3,
333
+ stride=stride,
334
+ padding=1,
335
+ bias=False))
336
+ bns.append(nn.BatchNorm2d(width))
337
+ self.convs = nn.ModuleList(convs)
338
+ self.bns = nn.ModuleList(bns)
339
+
340
+ self.conv3 = nn.Conv2d(width * scale,
341
+ planes * self.expansion,
342
+ kernel_size=kernel_size,
343
+ padding=(0,1),
344
+ bias=False)
345
+ self.bn3 = nn.BatchNorm2d(planes * self.expansion)
346
+ self.se = SELayer(planes * self.expansion, reduction=16)
347
+ self.relu = nn.ReLU(inplace=True)
348
+ if inplanes != planes:
349
+ self.downsample = True
350
+ self.conv_downsample = nn.Conv2d(in_channels=inplanes,
351
+ out_channels=planes,
352
+ padding=(0, 0),
353
+ kernel_size=(1, 1),
354
+ stride=1)
355
+
356
+ else:
357
+ self.downsample = False
358
+ self.stype = stype
359
+ self.scale = scale
360
+ self.width = width
361
+
362
+ def forward(self, x):
363
+ residual = x
364
+ out = self.conv1(x)
365
+ out = self.bn1(out)
366
+ out = self.relu(out)
367
+
368
+ spx = torch.split(out, self.width, 1)
369
+ for i in range(self.nums):
370
+ if i == 0 or self.stype == 'stage':
371
+ sp = spx[i]
372
+ else:
373
+ sp = sp + spx[i]
374
+ sp = self.convs[i](sp)
375
+ sp = self.relu(self.bns[i](sp))
376
+ if i == 0:
377
+ out = sp
378
+ else:
379
+ out = torch.cat((out, sp), 1)
380
+ if self.scale != 1 and self.stype == 'normal':
381
+ out = torch.cat((out, spx[self.nums]), 1)
382
+ elif self.scale != 1 and self.stype == 'stage':
383
+ out = torch.cat((out, self.pool(spx[self.nums])), 1)
384
+
385
+ out = self.conv3(out)
386
+ out = self.bn3(out)
387
+ out = self.se(out)
388
+
389
+ if self.downsample:
390
+ residual = self.conv_downsample(residual)
391
+
392
+ out += residual
393
+ out = self.relu(out)
394
+
395
+ return out
396
+
397
+ class Res2Net(nn.Module):
398
+ def __init__(self, block, layers, baseWidth=26, scale=4, m=0.35, num_classes=1000, loss='softmax', **kwargs):
399
+ self.inplanes = 16
400
+ super(Res2Net, self).__init__()
401
+ self.loss = loss
402
+ self.baseWidth = baseWidth
403
+ self.scale = scale
404
+ self.conv1 = nn.Sequential(nn.Conv2d(1, 16, 3, 1, 1, bias=False),
405
+ nn.BatchNorm2d(16), nn.ReLU(inplace=True),
406
+ nn.Conv2d(16, 16, 3, 1, 1, bias=False),
407
+ nn.BatchNorm2d(16), nn.ReLU(inplace=True),
408
+ nn.Conv2d(16, 16, 3, 1, 1, bias=False))
409
+ self.bn1 = nn.BatchNorm2d(16)
410
+ self.relu = nn.ReLU()
411
+ # self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
412
+ self.layer1 = self._make_layer(block, 16, layers[0])#64
413
+ self.layer2 = self._make_layer(block, 32, layers[1], stride=2)#128
414
+ self.layer3 = self._make_layer(block, 64, layers[2], stride=2)#256
415
+ self.layer4 = self._make_layer(block, 128, layers[3], stride=2)#512
416
+ self.avgpool = nn.AdaptiveAvgPool2d(1)
417
+ # self.stats_pooling = StatsPooling()
418
+
419
+ if self.loss == 'softmax':
420
+ # self.cls_layer = nn.Linear(2*8*128*block.expansion, num_classes)
421
+ self.cls_layer = nn.Linear(128*block.expansion, num_classes)
422
+ else:
423
+ raise NotImplementedError
424
+
425
+ for m in self.modules():
426
+ if isinstance(m, nn.Conv2d):
427
+ nn.init.kaiming_normal_(m.weight,
428
+ mode='fan_out',
429
+ nonlinearity='relu')
430
+ elif isinstance(m, nn.BatchNorm2d):
431
+ nn.init.constant_(m.weight, 1)
432
+ nn.init.constant_(m.bias, 0)
433
+
434
+ def _make_layer(self, block, planes, blocks, stride=1):
435
+ downsample = None
436
+ if stride != 1 or self.inplanes != planes * block.expansion:
437
+ downsample = nn.Sequential(
438
+ nn.AvgPool2d(kernel_size=stride,
439
+ stride=stride,
440
+ ceil_mode=True,
441
+ count_include_pad=False),
442
+ nn.Conv2d(self.inplanes,
443
+ planes * block.expansion,
444
+ kernel_size=1,
445
+ stride=1,
446
+ bias=False),
447
+ nn.BatchNorm2d(planes * block.expansion),
448
+ )
449
+
450
+ layers = []
451
+ layers.append(
452
+ block(self.inplanes,
453
+ planes,
454
+ stride,
455
+ downsample=downsample,
456
+ stype='stage',
457
+ baseWidth=self.baseWidth,
458
+ scale=self.scale))
459
+ self.inplanes = planes * block.expansion
460
+ for i in range(1, blocks):
461
+ layers.append(
462
+ block(self.inplanes,
463
+ planes,
464
+ baseWidth=self.baseWidth,
465
+ scale=self.scale))
466
+
467
+ return nn.Sequential(*layers)
468
+
469
+ def _forward(self, x):
470
+ x = x.unsqueeze(dim=1)
471
+ x = self.conv1(x)
472
+ x = self.bn1(x)
473
+ x = self.relu(x)
474
+ x = self.layer1(x)
475
+ x = self.layer2(x)
476
+ x = self.layer3(x)
477
+ x = self.layer4(x)
478
+ x = self.avgpool(x)
479
+ x = torch.flatten(x, 1)
480
+ x = self.cls_layer(x)
481
+
482
+ return F.log_softmax(x, dim=-1)
483
+
484
+ def extract(self, x):
485
+ x = self.conv1(x)
486
+ x = self.bn1(x)
487
+ x = self.relu(x)
488
+ x = self.layer1(x)
489
+ x = self.layer2(x)
490
+ x = self.layer3(x)
491
+ x = self.layer4(x)
492
+ x = self.avgpool(x)
493
+ x = torch.flatten(x, 1)
494
+ return x
495
+ # Allow for accessing forward method in a inherited class
496
+ forward = _forward
497
+
498
+ def se_res2net50_v1b_14w_8s(**kwargs):
499
+ """Constructs a Res2Net-50_v1b model.
500
+ Res2Net-50 refers to the Res2Net-50_v1b_26w_4s.
501
+ """
502
+ model = Res2Net(SEBottle2neck, [3, 4, 6, 3], baseWidth=14, scale=8, **kwargs)
503
+ return model
504
+
505
+
506
+ class CONV(nn.Module):
507
+ @staticmethod
508
+ def to_mel(hz):
509
+ return 2595 * np.log10(1 + hz / 700)
510
+
511
+ @staticmethod
512
+ def to_hz(mel):
513
+ return 700 * (10**(mel / 2595) - 1)
514
+
515
+ def __init__(self,
516
+ out_channels,
517
+ kernel_size,
518
+ sample_rate=16000,
519
+ in_channels=1,
520
+ stride=1,
521
+ padding=0,
522
+ dilation=1,
523
+ bias=False,
524
+ groups=1,
525
+ mask=False):
526
+ super().__init__()
527
+ if in_channels != 1:
528
+
529
+ msg = "SincConv only support one input channel (here, in_channels = {%i})" % (
530
+ in_channels)
531
+ raise ValueError(msg)
532
+ self.out_channels = out_channels
533
+ self.kernel_size = kernel_size
534
+ self.sample_rate = sample_rate
535
+
536
+ # Forcing the filters to be odd (i.e, perfectly symmetrics)
537
+ if kernel_size % 2 == 0:
538
+ self.kernel_size = self.kernel_size + 1
539
+ self.stride = stride
540
+ self.padding = padding
541
+ self.dilation = dilation
542
+ self.mask = mask
543
+ if bias:
544
+ raise ValueError('SincConv does not support bias.')
545
+ if groups > 1:
546
+ raise ValueError('SincConv does not support groups.')
547
+
548
+ NFFT = 512
549
+ f = int(self.sample_rate / 2) * np.linspace(0, 1, int(NFFT / 2) + 1)
550
+ fmel = self.to_mel(f)
551
+ fmelmax = np.max(fmel)
552
+ fmelmin = np.min(fmel)
553
+ filbandwidthsmel = np.linspace(fmelmin, fmelmax, self.out_channels + 1)
554
+ filbandwidthsf = self.to_hz(filbandwidthsmel)
555
+
556
+ self.mel = filbandwidthsf
557
+ self.hsupp = torch.arange(-(self.kernel_size - 1) / 2,
558
+ (self.kernel_size - 1) / 2 + 1)
559
+ self.band_pass = torch.zeros(self.out_channels, self.kernel_size)
560
+ for i in range(len(self.mel) - 1):
561
+ fmin = self.mel[i]
562
+ fmax = self.mel[i + 1]
563
+ hHigh = (2*fmax/self.sample_rate) * \
564
+ np.sinc(2*fmax*self.hsupp/self.sample_rate)
565
+ hLow = (2*fmin/self.sample_rate) * \
566
+ np.sinc(2*fmin*self.hsupp/self.sample_rate)
567
+ hideal = hHigh - hLow
568
+
569
+ self.band_pass[i, :] = Tensor(np.hamming(
570
+ self.kernel_size)) * Tensor(hideal)
571
+
572
+ def forward(self, x, mask=False):
573
+ band_pass_filter = self.band_pass.clone().to(x.device)
574
+ if mask:
575
+ A = np.random.uniform(0, 20)
576
+ A = int(A)
577
+ A0 = random.randint(0, band_pass_filter.shape[0] - A)
578
+ band_pass_filter[A0:A0 + A, :] = 0
579
+ else:
580
+ band_pass_filter = band_pass_filter
581
+
582
+ self.filters = (band_pass_filter).view(self.out_channels, 1,
583
+ self.kernel_size)
584
+
585
+ return F.conv1d(x,
586
+ self.filters,
587
+ stride=self.stride,
588
+ padding=self.padding,
589
+ dilation=self.dilation,
590
+ bias=None,
591
+ groups=1)
592
+
593
+
594
+ class My_Residual_block(nn.Module):
595
+ def __init__(self, nb_filts, first=False, conv1=[2, 3, 1, 1, 1, 1], conv2=[2, 3, 0, 1, 1, 3], conv3=[1, 3, 0, 1, 1, 3], pool=(1, 3)):
596
+ super().__init__()
597
+ self.first = first
598
+
599
+ if not self.first:
600
+ self.bn1 = nn.BatchNorm2d(num_features=nb_filts[0])
601
+ self.conv1 = nn.Conv2d(in_channels=nb_filts[0],
602
+ out_channels=nb_filts[1],
603
+ kernel_size=(conv1[0], conv1[1]),
604
+ padding=(conv1[2], conv1[3]),
605
+ stride=(conv1[4], conv1[5]))
606
+ self.selu = nn.SELU(inplace=True)
607
+
608
+ self.bn2 = nn.BatchNorm2d(num_features=nb_filts[1])
609
+ self.conv2 = nn.Conv2d(in_channels=nb_filts[1],
610
+ out_channels=nb_filts[1],
611
+ kernel_size=(conv2[0], conv2[1]),
612
+ padding=(conv2[2], conv2[3]),
613
+ stride=(conv2[4], conv2[5]))
614
+
615
+ self.downsample = True
616
+ self.conv_downsample = nn.Conv2d(in_channels=nb_filts[0],
617
+ out_channels=nb_filts[1],
618
+ kernel_size=(conv3[0], conv3[1]),
619
+ padding=(conv3[2], conv3[3]),
620
+ stride=(conv3[4], conv3[5]))
621
+
622
+ # self.mp = nn.MaxPool2d((1,4))
623
+ self.mp = nn.MaxPool2d((pool[0], pool[1]))
624
+
625
+ def forward(self, x):
626
+ identity = x
627
+ if not self.first:
628
+ out = self.bn1(x)
629
+ out = self.selu(out)
630
+ else:
631
+ out = x
632
+ out = self.conv1(x)
633
+
634
+ out = self.bn2(out)
635
+ out = self.selu(out)
636
+ out = self.conv2(out)
637
+
638
+ if self.downsample:
639
+ identity = self.conv_downsample(identity)
640
+
641
+ out += identity
642
+ out = self.mp(out)
643
+ return out
644
+
645
+
646
+ class My_SERes2Net_block(nn.Module):
647
+ def __init__(self, nb_filts, first=False, conv1=[2, 3, 1, 1, 1, 1], conv2=[3, 3, 1, 1, 1, 3], conv3=[1, 3, 0, 1, 1, 3], pool=(1, 3), radix=2, groups=2):
648
+ super().__init__()
649
+ self.first = first
650
+
651
+ if not self.first:
652
+ self.bn1 = nn.BatchNorm2d(num_features=nb_filts[0])
653
+
654
+ self.conv1 = SEBottle2neck(inplanes=nb_filts[0],
655
+ planes=nb_filts[1], kernel_size=(conv1[0], conv1[1]))
656
+ self.selu = nn.SELU(inplace=True)
657
+
658
+ self.bn2 = nn.BatchNorm2d(num_features=nb_filts[1])
659
+ self.conv2 = nn.Conv2d(in_channels=nb_filts[1],
660
+ out_channels=nb_filts[1],
661
+ kernel_size=(conv2[0], conv2[1]),
662
+ padding=(conv2[2], conv2[3]),
663
+ stride=(conv2[4], conv2[5]))
664
+
665
+ self.downsample = True
666
+ self.conv_downsample = nn.Conv2d(in_channels=nb_filts[0],
667
+ out_channels=nb_filts[1],
668
+ kernel_size=(conv3[0], conv3[1]),
669
+ padding=(conv3[2], conv3[3]),
670
+ stride=(conv3[4], conv3[5]))
671
+
672
+ self.mp = nn.MaxPool2d((pool[0], pool[1]))
673
+
674
+ def forward(self, x):
675
+ identity = x
676
+ if not self.first:
677
+ out = self.bn1(x)
678
+ out = self.selu(out)
679
+ else:
680
+ out = x
681
+
682
+ out = self.conv1(x)
683
+ out = self.bn2(out)
684
+ out = self.selu(out)
685
+ out = self.conv2(out)
686
+ if self.downsample:
687
+ identity = self.conv_downsample(identity)
688
+
689
+ out += identity
690
+ out = self.mp(out)
691
+ return out
692
+
693
+
694
+ class Attention(Module):
695
+ """
696
+ Obtained from timm: github.com:rwightman/pytorch-image-models
697
+ """
698
+
699
+ def __init__(self, dim, num_heads=8, attention_dropout=0.1, projection_dropout=0.1):
700
+ super().__init__()
701
+ self.num_heads = num_heads
702
+ head_dim = dim // self.num_heads
703
+ self.scale = head_dim ** -0.5
704
+
705
+ self.qkv = Linear(dim, dim * 3, bias=False)
706
+ self.attn_drop = Dropout(attention_dropout)
707
+ self.proj = Linear(dim, dim)
708
+ self.proj_drop = Dropout(projection_dropout)
709
+
710
+ def forward(self, x):
711
+ B, N, C = x.shape
712
+ qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C //
713
+ self.num_heads).permute(2, 0, 3, 1, 4)
714
+ q, k, v = qkv[0], qkv[1], qkv[2]
715
+
716
+ attn = (q @ k.transpose(-2, -1)) * self.scale
717
+ attn = attn.softmax(dim=-1)
718
+ attn = self.attn_drop(attn)
719
+
720
+ x = (attn @ v).transpose(1, 2).reshape(B, N, C)
721
+ x = self.proj(x)
722
+ x = self.proj_drop(x)
723
+ return x
724
+
725
+
726
+ class SimpleRelativeAttention(nn.Module):
727
+ # we implement this relative position embedding here., this is not used in our experiments
728
+ def __init__(self, dim, seq_length, num_heads=8, qkv_bias=True, qk_scale=None, attn_drop=0.1, proj_drop=0.1):
729
+ super().__init__()
730
+ self.dim = dim
731
+ self.length = seq_length
732
+ self.num_heads = num_heads
733
+ head_dim = dim//num_heads
734
+ self.scale = qk_scale or head_dim**-0.5
735
+ self.relative_position_table = nn.Parameter(
736
+ torch.zeros(size=(seq_length*2-1, num_heads)))
737
+ coords = torch.arange(seq_length)
738
+ relative_coords = coords[:, None]-coords[None, :]
739
+ relative_coords = relative_coords+seq_length-1
740
+ self.register_buffer('relative_index', relative_coords)
741
+ self.qkv = nn.Linear(dim, dim*3)
742
+ self.attn_drop = nn.Dropout(attn_drop)
743
+ self.proj = nn.Linear(dim, dim)
744
+ self.proj_drop = nn.Dropout(proj_drop)
745
+
746
+ trunc_normal_(self.relative_position_table, std=0.02)
747
+
748
+ def forward(self, x):
749
+ B, N, C = x.shape
750
+ qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C //
751
+ self.num_heads).permute(2, 0, 3, 1, 4)
752
+ q, k, v = qkv[0], qkv[1], qkv[2]
753
+ q = q*self.scale
754
+ attn = torch.einsum('bhqe,bhke->bhqk', q, k)
755
+ relative_position_bias = self.relative_position_table[self.relative_index.reshape(-1)].reshape(
756
+ self.length, self.length, self.num_heads
757
+ )
758
+ relative_position_bias = relative_position_bias.permute(
759
+ 2, 0, 1).contiguous()
760
+ attn = attn+relative_position_bias.unsqueeze(0)
761
+ attn = attn.softmax(-1)
762
+ attn = self.attn_drop(attn)
763
+ x = torch.einsum('bnqk,bnqe->bnqe', attn,
764
+ v).transpose(1, 2).reshape(B, N, C)
765
+ x = self.proj_drop(self.proj(x))
766
+ return x
767
+
768
+
769
+ class TransformerEncoderLayer(Module):
770
+ """
771
+ Inspired by torch.nn.TransformerEncoderLayer and timm.
772
+ """
773
+
774
+ def __init__(self, d_model, nhead, atten=Attention, dim_feedforward=2048, dropout=0.1,
775
+ attention_dropout=0.1, drop_path_rate=0.1):
776
+ super(TransformerEncoderLayer, self).__init__()
777
+ self.pre_norm = LayerNorm(d_model)
778
+ self.self_attn = atten(dim=d_model, num_heads=nhead,
779
+ attention_dropout=attention_dropout, projection_dropout=dropout)
780
+
781
+ self.linear1 = Linear(d_model, dim_feedforward)
782
+ self.dropout1 = Dropout(dropout)
783
+ self.norm1 = LayerNorm(d_model)
784
+ self.linear2 = Linear(dim_feedforward, d_model)
785
+ self.dropout2 = Dropout(dropout)
786
+
787
+ self.drop_path = DropPath(
788
+ drop_path_rate) if drop_path_rate > 0 else Identity()
789
+
790
+ self.activation = F.gelu
791
+
792
+ def forward(self, src: torch.Tensor, *args, **kwargs) -> torch.Tensor:
793
+ src = src + self.drop_path(self.self_attn(self.pre_norm(src)))
794
+ src = self.norm1(src)
795
+ src2 = self.linear2(self.dropout1(self.activation(self.linear1(src))))
796
+ src = src + self.drop_path(self.dropout2(src2))
797
+ return src
798
+
799
+
800
+ class TransformerClassifier(Module):
801
+ """
802
+ Adopted from https://github.com/SHI-Labs/Compact-Transformers.git
803
+ """
804
+
805
+ def __init__(self,
806
+ embedding_dim=768,
807
+ num_classes=1000,
808
+ num_layers=12,
809
+ num_heads=12,
810
+ mlp_ratio=4.0,
811
+ dropout_rate=0.1,
812
+ attention_dropout=0.1,
813
+ stochastic_depth_rate=0.1,
814
+ positional_embedding='sine',
815
+ sequence_length=10000,
816
+ *args, **kwargs):
817
+ super().__init__()
818
+ positional_embedding = positional_embedding if \
819
+ positional_embedding in ['sine', 'learnable', 'none'] else 'sine'
820
+ dim_feedforward = int(embedding_dim * mlp_ratio)
821
+ self.embedding_dim = embedding_dim
822
+ self.sequence_length = sequence_length
823
+
824
+ assert sequence_length is not None or positional_embedding == 'none'
825
+
826
+ if positional_embedding != 'none':
827
+ if positional_embedding == 'learnable':
828
+ self.positional_emb = Parameter(torch.zeros(1, sequence_length, embedding_dim),
829
+ requires_grad=True)
830
+ init.trunc_normal_(self.positional_emb, std=0.2)
831
+ else:
832
+ print('here!!! sinusoidal_embedding')
833
+ self.positional_emb = Parameter(self.sinusoidal_embedding(sequence_length, embedding_dim),
834
+ requires_grad=False)
835
+ else:
836
+ self.positional_emb = None
837
+
838
+ self.dropout = Dropout(p=dropout_rate)
839
+ dpr = [x.item() for x in torch.linspace(
840
+ 0, stochastic_depth_rate, num_layers)]
841
+ self.blocks = ModuleList([
842
+ TransformerEncoderLayer(d_model=embedding_dim, nhead=num_heads,
843
+ dim_feedforward=dim_feedforward, dropout=dropout_rate,
844
+ attention_dropout=attention_dropout, drop_path_rate=dpr[i])
845
+ for i in range(num_layers)])
846
+ self.norm = LayerNorm(embedding_dim)
847
+ self.flattener = nn.Flatten(2, 3)
848
+ self.attention_pool = Linear(self.embedding_dim, 1)
849
+ self.fc = Linear(embedding_dim, num_classes)
850
+ self.apply(self.init_weight)
851
+
852
+ def forward(self, x):
853
+ x = torch.transpose(x,-1,-2)
854
+ seq_len = x.size(1)
855
+ x += self.positional_emb[:, :seq_len, :]
856
+
857
+ x = self.dropout(x)
858
+ for blk in self.blocks:
859
+ x = blk(x)
860
+ x = self.norm(x)
861
+
862
+ feature = torch.matmul(F.softmax(self.attention_pool(
863
+ x), dim=1).transpose(-1, -2), x).squeeze(-2)
864
+ logits = self.fc(feature)
865
+
866
+ return logits, feature
867
+
868
+
869
+ @staticmethod
870
+ def init_weight(m):
871
+ if isinstance(m, Linear):
872
+ init.trunc_normal_(m.weight, std=.02)
873
+ if isinstance(m, Linear) and m.bias is not None:
874
+ init.constant_(m.bias, 0)
875
+ elif isinstance(m, LayerNorm):
876
+ init.constant_(m.bias, 0)
877
+ init.constant_(m.weight, 1.0)
878
+
879
+ @staticmethod
880
+ def sinusoidal_embedding(n_channels, dim):
881
+ pe = torch.FloatTensor([[p / (10000 ** (2 * (i // 2) / dim)) for i in range(dim)]
882
+ for p in range(n_channels)])
883
+ pe[:, 0::2] = torch.sin(pe[:, 0::2])
884
+ pe[:, 1::2] = torch.cos(pe[:, 1::2])
885
+ return pe.unsqueeze(0)
886
+
887
+ class SE_Rawformer_front(nn.Module):
888
+ def __init__(self, conv1 = [2,3,1,1,1,1],conv2 = [3,3,1,1,1,2],conv3 = [1,3,0,1,1,2]):
889
+ super().__init__()
890
+ filts = [70, [1, 32], [32, 32], [32, 64], [64, 64]]
891
+ self.conv_time = CONV(out_channels=filts[0],
892
+ kernel_size=128,
893
+ in_channels=1) # 70 129
894
+ self.first_bn = nn.BatchNorm2d(num_features=1)
895
+ self.drop = nn.Dropout(0.5, inplace=True)
896
+ self.drop_way = nn.Dropout(0.2, inplace=True)
897
+ self.selu = nn.SELU(inplace=True)
898
+
899
+ self.encoder = nn.Sequential(
900
+ nn.Sequential(My_Residual_block(nb_filts=filts[1], conv1 = conv1,conv2 = [2,3,0,1,1,2],conv3 = conv3,first=True)),
901
+ nn.Sequential(My_SERes2Net_block(nb_filts=filts[2], conv1 = conv1,conv2 = conv2,conv3 = conv3)),
902
+ nn.Sequential(My_SERes2Net_block(nb_filts=filts[3], conv1 = conv1,conv2 = conv2,conv3 = conv3)),
903
+ nn.Sequential(My_SERes2Net_block(nb_filts=filts[4], conv1 = conv1,conv2 = conv2,conv3 = conv3)))
904
+
905
+ def forward(self, x, Freq_aug=False):
906
+ x = self.conv_time(x, mask=Freq_aug)
907
+ x = x.unsqueeze(dim=1)
908
+ x = F.max_pool2d(torch.abs(x), (3, 3))
909
+ x = self.first_bn(x)
910
+ x = self.selu(x)
911
+
912
+ encoder = self.encoder(x)
913
+ return encoder
914
+
915
+
916
+ class SafeEar(nn.Module):
917
+ def __init__(self,front, *args, **kwargs):
918
+ super().__init__()
919
+ # self.front = front
920
+ self.bottleneck = nn.Sequential(
921
+ nn.Conv1d(kwargs["embedding_dim"]*7, kwargs["embedding_dim"], kernel_size=1),
922
+ nn.BatchNorm1d(kwargs["embedding_dim"])
923
+ )
924
+ self.classifier = TransformerClassifier(*args, **kwargs)
925
+
926
+ def forward(self, encoder):
927
+ encoder = self.bottleneck(torch.cat(encoder, dim=1))
928
+ batch_size, feature_dim, frame_num = encoder.size()
929
+
930
+ for i in range(0, frame_num, 50):
931
+ encoder[:, :, i:i+50] = torch.flip(encoder[:, :, i:i+50], dims=[2])
932
+
933
+ logits, feature = self.classifier(encoder)
934
+
935
+ return logits, feature
936
+
937
+
938
+ class SafeEar1s(nn.Module):
939
+ def __init__(self,front, *args, **kwargs):
940
+ super().__init__()
941
+ # self.front = front
942
+ self.bottleneck = nn.Sequential(
943
+ nn.Conv1d(kwargs["embedding_dim"]*7, kwargs["embedding_dim"], kernel_size=1),
944
+ nn.BatchNorm1d(kwargs["embedding_dim"])
945
+ )
946
+ self.classifier = TransformerClassifier(*args, **kwargs)
947
+
948
+ def forward(self, encoder):
949
+ encoder = self.bottleneck(torch.cat(encoder, dim=1))
950
+
951
+ batch_size, feature_dim, frame_num = encoder.size()
952
+ for i in range(0, frame_num, 50):
953
+ end = i+50 if i+50 <= frame_num else frame_num
954
+ indices = torch.randperm(end - i)
955
+ encoder[:, :, i:end] = encoder[:, :, i+indices]
956
+
957
+ logits, feature = self.classifier(encoder)
958
+
959
+ return logits, feature
clean/audio/safeear/safeear/trainer/safeear_trainer.py ADDED
@@ -0,0 +1,189 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ import torch
3
+ import pytorch_lightning as pl
4
+ from ..losses.loss import compute_eer
5
+ import numpy as np
6
+ import warnings
7
+ warnings.filterwarnings("ignore")
8
+
9
+ def get_input(x):
10
+ x = x.to(memory_format=torch.contiguous_format)
11
+ return x.float()
12
+
13
+ class SafeEarTrainer(pl.LightningModule):
14
+ def __init__(
15
+ self,
16
+ decouple_model,
17
+ detect_model,
18
+ lr_raw_former,
19
+ save_score_path
20
+ ) -> None:
21
+ super().__init__()
22
+
23
+ self.decouple_model = decouple_model
24
+ self.detect_model = detect_model
25
+ self.lr_raw_former = lr_raw_former
26
+ self.save_score_path = save_score_path
27
+
28
+ self.detect_loss = torch.nn.BCELoss()
29
+
30
+ self.automatic_optimization = False
31
+
32
+ self.val_index_loader = []
33
+ self.val_score_loader = []
34
+ self.eval_index_loader = []
35
+ self.eval_score_loader = []
36
+ self.eval_filename_loader = []
37
+ self.default_monitor = "val_eer"
38
+
39
+ def forward(self, batch, is_train=True):
40
+ if is_train:
41
+ x, feat, target = batch
42
+ else:
43
+ if len(batch) == 4:
44
+ x, feat, target, audio_path = batch
45
+ else:
46
+ x, feat, target = batch
47
+ audio_path = None
48
+ x_wav = get_input(x)
49
+ with torch.no_grad():
50
+ self.decouple_model.eval()
51
+ G_x, commit_loss, last_layer, acoustic_tokens = self.decouple_model(x_wav, layers=[0,1,2,3,4,5,6,7])
52
+ raw_logits, raw_feature = self.detect_model(acoustic_tokens)
53
+
54
+ if is_train:
55
+ onehot_target = torch.eye(2).to(self.device)[target, :]
56
+ raw_logits = torch.softmax(raw_logits, dim=-1)
57
+ raw_former_loss_ = self.detect_loss(raw_logits,onehot_target)
58
+ return raw_former_loss_, raw_logits, target
59
+ else:
60
+ raw_logits = torch.softmax(raw_logits, dim=-1)[:, 0]
61
+ raw_former_loss_ = 0
62
+ return audio_path, raw_former_loss_, raw_logits, target
63
+
64
+ def training_step(self, batch, batch_idx):
65
+ raw_opt = self.optimizers()
66
+
67
+ raw_former_loss_, raw_logits, target = self(batch, is_train=True)
68
+ raw_opt.zero_grad()
69
+ self.manual_backward(raw_former_loss_)
70
+ raw_opt.step()
71
+
72
+ self.log_dict(
73
+ {
74
+ 'train_loss': raw_former_loss_
75
+ },
76
+ on_step=True,
77
+ on_epoch=True,
78
+ prog_bar=True,
79
+ sync_dist=True,
80
+ logger=True)
81
+
82
+ def validation_step(self, batch, batch_idx):
83
+ _, raw_former_loss_, raw_logits, target = self(batch, is_train=False)
84
+
85
+ self.val_index_loader.append(target)
86
+ self.val_score_loader.append(raw_logits)
87
+
88
+ self.log_dict(
89
+ {
90
+ 'val_loss': raw_former_loss_,
91
+ },
92
+ on_epoch=True,
93
+ prog_bar=True,
94
+ sync_dist=True,
95
+ logger=True)
96
+
97
+ def on_validation_epoch_end(self):
98
+ all_index = self.all_gather(torch.cat(self.val_index_loader, dim=0)).view(-1).cpu().numpy()
99
+ all_score = self.all_gather(torch.cat(self.val_score_loader, dim=0)).view(-1).cpu().numpy()
100
+ val_eer = compute_eer(all_score[all_index == 0], all_score[all_index == 1])[0]
101
+ other_val_eer = compute_eer(-all_score[all_index == 0], -all_score[all_index == 1])[0]
102
+ val_eer = min(val_eer, other_val_eer)
103
+ self.log_dict(
104
+ {
105
+ "val_eer": val_eer,
106
+ },
107
+ sync_dist=True,
108
+ on_epoch=True,
109
+ prog_bar=True,
110
+ logger=True)
111
+
112
+ self.val_index_loader.clear() # free memory
113
+ self.val_score_loader.clear() # free memory
114
+
115
+ self.log_dict(
116
+ {
117
+ "lr": self.optimizers().param_groups[0]['lr'],
118
+ },
119
+ sync_dist=True,
120
+ on_epoch=True,
121
+ prog_bar=False,
122
+ logger=True
123
+ )
124
+
125
+ adjust_learning_rate(self.optimizers(), self.current_epoch, self.lr_raw_former, self.trainer.max_epochs*0.1, self.trainer.max_epochs)
126
+
127
+ def test_step(self, batch, batch_idx):
128
+
129
+ audio_path, raw_former_loss_, raw_logits, target = self(batch, is_train=False)
130
+
131
+ self.eval_index_loader.append(target)
132
+ self.eval_score_loader.append(raw_logits)
133
+ self.eval_filename_loader.append(audio_path)
134
+ self.log_dict(
135
+ {
136
+ 'val_loss_rawformer': raw_former_loss_,
137
+ },
138
+ on_epoch=True,
139
+ prog_bar=True,
140
+ sync_dist=True,
141
+ logger=True)
142
+
143
+ def on_test_epoch_end(self):
144
+
145
+ string_list = [list(item) for item in self.eval_filename_loader]
146
+
147
+ all_filename = np.array(string_list)
148
+ all_filename = all_filename.reshape(-1, 1)
149
+
150
+
151
+ all_index = self.all_gather(torch.cat(self.eval_index_loader, dim=0)).view(-1).cpu().numpy()
152
+ all_score = self.all_gather(torch.cat(self.eval_score_loader, dim=0)).view(-1).cpu().numpy()
153
+
154
+ # gpu_id = torch.cuda.current_device()
155
+
156
+ data_to_write = zip(all_filename, all_score,all_index)
157
+ csv_filename = self.save_score_path + '/score.csv'
158
+ eval_eer = compute_eer(all_score[all_index == 0], all_score[all_index == 1])[0]
159
+ other_eval_eer = compute_eer(-all_score[all_index == 0], -all_score[all_index == 1])[0]
160
+ eval_eer = min(eval_eer, other_eval_eer)
161
+
162
+ self.log_dict(
163
+ {
164
+ "test_eer": eval_eer,
165
+ },
166
+ sync_dist=True,
167
+ on_epoch=True,
168
+ prog_bar=True,
169
+ logger=True)
170
+
171
+ self.eval_index_loader.clear() # free memory
172
+ self.eval_score_loader.clear() # free memory
173
+ self.eval_filename_loader.clear() # free memory
174
+
175
+ def configure_optimizers(self):
176
+ optimizer_rawformer = torch.optim.AdamW(self.detect_model.parameters(), lr=self.lr_raw_former, weight_decay=1e-4)
177
+
178
+ return [optimizer_rawformer]
179
+
180
+ def adjust_learning_rate(optimizer, epoch, lr, warmup, epochs=100):
181
+ lr = lr
182
+ if epoch < warmup:
183
+ lr = lr / (warmup - epoch)
184
+ else:
185
+ lr *= 0.5 * (1. + math.cos(math.pi *
186
+ (epoch - warmup) / (epochs - warmup)))
187
+
188
+ for param_group in optimizer.param_groups:
189
+ param_group['lr'] = lr
clean/audio/safeear/safeear/utils/dump_hubert_feature.py ADDED
@@ -0,0 +1,108 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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 tqdm
10
+ sys.path.append('../../fairseq_ours/') # we recommend an abosulte path here.
11
+ import fairseq
12
+ import soundfile as sf
13
+ import torch
14
+ import torch.nn.functional as F
15
+ from npy_append_array import NpyAppendArray
16
+ from pathlib import Path
17
+ # from feature_utils import get_path_iterator, dump_feature
18
+
19
+
20
+ logging.basicConfig(
21
+ format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
22
+ datefmt="%Y-%m-%d %H:%M:%S",
23
+ level=os.environ.get("LOGLEVEL", "INFO").upper(),
24
+ stream=sys.stdout,
25
+ )
26
+ logger = logging.getLogger("dump_hubert_feature")
27
+
28
+
29
+ class HubertFeatureReader(object):
30
+ def __init__(self, ckpt_path, layer, max_chunk=1600000):
31
+ (
32
+ model,
33
+ cfg,
34
+ task,
35
+ ) = fairseq.checkpoint_utils.load_model_ensemble_and_task([ckpt_path])
36
+ self.model = model[0].eval().cuda()
37
+ self.task = task
38
+ self.layer = layer
39
+ self.max_chunk = max_chunk
40
+ logger.info(f"TASK CONFIG:\n{self.task.cfg}")
41
+ logger.info(f" max_chunk = {self.max_chunk}")
42
+
43
+ def read_audio(self, path, ref_len=None):
44
+ wav, sr = sf.read(path)
45
+ assert sr == self.task.cfg.sample_rate, sr
46
+ if wav.ndim == 2:
47
+ wav = wav.mean(-1)
48
+ assert wav.ndim == 1, wav.ndim
49
+ if ref_len is not None and abs(ref_len - len(wav)) > 160:
50
+ logging.warning(f"ref {ref_len} != read {len(wav)} ({path})")
51
+ return wav
52
+
53
+ def get_feats(self, path, ref_len=None):
54
+ x = self.read_audio(path, ref_len)
55
+ with torch.no_grad():
56
+ x = torch.from_numpy(x).float().cuda()
57
+ if self.task.cfg.normalize:
58
+ x = F.layer_norm(x, x.shape)
59
+ x = x.view(1, -1)
60
+
61
+ feat = []
62
+ for start in range(0, x.size(1), self.max_chunk):
63
+ x_chunk = x[:, start: start + self.max_chunk]
64
+ feat_chunk, _, _ = self.model.extract_features(
65
+ source=x_chunk,
66
+ padding_mask=None,
67
+ mask=False,
68
+ output_layer=self.layer,
69
+ )
70
+ feat.append(feat_chunk)
71
+ return torch.cat(feat, 1).squeeze(0)
72
+
73
+ def dump_feature(reader,audio_dir,save_dir):
74
+ save_dir = Path(save_dir)
75
+ audio_dir = Path(audio_dir)
76
+
77
+ audio_files = list(audio_dir.glob("**/*.flac"))
78
+
79
+ for audio_file in tqdm.tqdm(audio_files):
80
+ releative_path = audio_file.relative_to(audio_dir).with_suffix(".npy")
81
+ save_path = save_dir / releative_path
82
+ # import pdb; pdb.set_trace()
83
+ if not save_path.parent.exists():
84
+ save_path.parent.mkdir(parents=True)
85
+
86
+ feat_f = NpyAppendArray(save_path)
87
+ feat = reader.get_feats(audio_file)
88
+ feat_f.append(feat.cpu().numpy())
89
+ logger.info("finished successfully")
90
+
91
+ def main(audio_dir, save_dir, ckpt_path, layer, max_chunk):
92
+ reader = HubertFeatureReader(ckpt_path, layer, max_chunk)
93
+ dump_feature(reader, audio_dir, save_dir)
94
+
95
+ if __name__ == "__main__":
96
+ import argparse
97
+
98
+ parser = argparse.ArgumentParser()
99
+ parser.add_argument("--audio_dir", default="./datasets/ASVSpoof2021/ASVspoof2021_LA_eval/flac")
100
+ parser.add_argument("--save_dir", default="./datasets/ASVSpoof2021/ASVspoof2021_LA_eval/Hubert_L9")
101
+
102
+ parser.add_argument("--ckpt_path", default="./model_zoo/hubert/hubert_base_ls960.pt")
103
+ parser.add_argument("--layer", type=int, default=9)
104
+ parser.add_argument("--max_chunk", type=int, default=1600000)
105
+ args = parser.parse_args()
106
+ logger.info(args)
107
+
108
+ main(**vars(args))
clean/audio/safeear/test.py ADDED
@@ -0,0 +1,78 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import importlib
2
+ import json
3
+ import os
4
+ from typing import Any, Dict, List, Optional, Tuple
5
+ import argparse
6
+ import pytorch_lightning as pl
7
+ import torch
8
+ import hydra
9
+
10
+ torch.set_float32_matmul_precision("high")
11
+
12
+ from pytorch_lightning import Callback, LightningDataModule, LightningModule, Trainer
13
+ from pytorch_lightning.strategies.ddp import DDPStrategy
14
+ from omegaconf import DictConfig
15
+ from omegaconf import OmegaConf
16
+ from pytorch_lightning.utilities import rank_zero_only
17
+
18
+ @rank_zero_only
19
+ def print_only(message: str):
20
+ """Prints a message only on rank 0."""
21
+ print(message)
22
+
23
+ def train(cfg: DictConfig, args) -> Tuple[Dict[str, Any], Dict[str, Any]]:
24
+
25
+ # instantiate datamodule
26
+ print_only(f"Instantiating datamodule <{cfg.datamodule._target_}>")
27
+ datamodule: LightningDataModule = hydra.utils.instantiate(cfg.datamodule)
28
+
29
+ # instantiate decouple model
30
+ print_only(f"Instantiating decouple model <{cfg.decouple_model._target_}>")
31
+ decouple_model: torch.nn.Module = hydra.utils.instantiate(cfg.decouple_model)
32
+ decouple_model.load_state_dict(torch.load(cfg.speechtokenizer_path))
33
+ # import pdb; pdb.set_trace()
34
+
35
+ # instantiate detect model
36
+ print(f"Instantiating detect model <{cfg.detect_model._target_}>")
37
+ detect_model: torch.nn.Module = hydra.utils.instantiate(cfg.detect_model)
38
+ # import pdb; pdb.set_trace()
39
+
40
+ # instantiate system
41
+ print_only(f"Instantiating system <{cfg.system._target_}>")
42
+ system: LightningModule = hydra.utils.instantiate(
43
+ cfg.system,
44
+ decouple_model=decouple_model,
45
+ detect_model=detect_model,
46
+ )
47
+
48
+ # instantiate trainer
49
+ print_only(f"Instantiating trainer <{cfg.trainer._target_}>")
50
+ trainer: Trainer = hydra.utils.instantiate(
51
+ cfg.trainer,
52
+ strategy=DDPStrategy(find_unused_parameters=True),
53
+ )
54
+
55
+ trainer.test(system, datamodule=datamodule, ckpt_path=args.ckpt_path)
56
+
57
+ if __name__ == "__main__":
58
+
59
+ parser = argparse.ArgumentParser()
60
+ parser.add_argument(
61
+ "--conf_dir",
62
+ default="local/conf.yml",
63
+ help="Full path to save best validation model",
64
+ )
65
+ parser.add_argument(
66
+ "--ckpt_path",
67
+ help="Full path to save best validation model",
68
+ )
69
+
70
+ args = parser.parse_args()
71
+ cfg = OmegaConf.load(args.conf_dir)
72
+
73
+ os.makedirs(os.path.join(cfg.exp.dir, cfg.exp.name), exist_ok=True)
74
+ # 保存配置到新的文件
75
+ OmegaConf.save(cfg, os.path.join(cfg.exp.dir, cfg.exp.name, "config.yaml"))
76
+
77
+ train(cfg, args)
78
+
clean/audio/safeear/train.py ADDED
@@ -0,0 +1,102 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import importlib
2
+ import json
3
+ import os
4
+ import warnings
5
+ warnings.filterwarnings("ignore")
6
+ from typing import Any, Dict, List, Optional, Tuple
7
+ import argparse
8
+ import pytorch_lightning as pl
9
+ import torch
10
+ import hydra
11
+
12
+ torch.set_float32_matmul_precision("high")
13
+
14
+ from pytorch_lightning import Callback, LightningDataModule, LightningModule, Trainer
15
+ from pytorch_lightning.strategies.ddp import DDPStrategy
16
+ from omegaconf import DictConfig
17
+ from omegaconf import OmegaConf
18
+ from pytorch_lightning.utilities import rank_zero_only
19
+
20
+ @rank_zero_only
21
+ def print_only(message: str):
22
+ """Prints a message only on rank 0."""
23
+ print(message)
24
+
25
+ def train(cfg: DictConfig) -> Tuple[Dict[str, Any], Dict[str, Any]]:
26
+
27
+ # instantiate datamodule
28
+ print_only(f"Instantiating datamodule <{cfg.datamodule._target_}>")
29
+ datamodule: LightningDataModule = hydra.utils.instantiate(cfg.datamodule)
30
+ datamodule.setup()
31
+
32
+ # instantiate decouple model
33
+ print_only(f"Instantiating decouple model <{cfg.decouple_model._target_}>")
34
+ decouple_model: torch.nn.Module = hydra.utils.instantiate(cfg.decouple_model)
35
+ decouple_model.load_state_dict(torch.load(cfg.speechtokenizer_path))
36
+ # import pdb; pdb.set_trace()
37
+
38
+ # instantiate detect model
39
+ print(f"Instantiating detect model <{cfg.detect_model._target_}>")
40
+ detect_model: torch.nn.Module = hydra.utils.instantiate(cfg.detect_model)
41
+ # import pdb; pdb.set_trace()
42
+
43
+ # instantiate system
44
+ print_only(f"Instantiating system <{cfg.system._target_}>")
45
+ system: LightningModule = hydra.utils.instantiate(
46
+ cfg.system,
47
+ decouple_model=decouple_model,
48
+ detect_model=detect_model,
49
+ )
50
+ # instantiate callbacks
51
+ callbacks: List[Callback] = []
52
+ if cfg.get("early_stopping"):
53
+ print_only(f"Instantiating early_stopping <{cfg.early_stopping._target_}>")
54
+ callbacks.append(hydra.utils.instantiate(cfg.early_stopping))
55
+ if cfg.get("checkpoint"):
56
+ print_only(f"Instantiating checkpoint <{cfg.checkpoint._target_}>")
57
+ checkpoint: pl.callbacks.ModelCheckpoint = hydra.utils.instantiate(cfg.checkpoint)
58
+ callbacks.append(checkpoint)
59
+
60
+ # instantiate logger
61
+ print_only(f"Instantiating logger <{cfg.logger._target_}>")
62
+ os.makedirs(os.path.join(cfg.exp.dir, cfg.exp.name, "logs"), exist_ok=True)
63
+ logger = hydra.utils.instantiate(cfg.logger)
64
+
65
+ # instantiate trainer
66
+ print_only(f"Instantiating trainer <{cfg.trainer._target_}>")
67
+ trainer: Trainer = hydra.utils.instantiate(
68
+ cfg.trainer,
69
+ callbacks=callbacks,
70
+ logger=logger,
71
+ strategy=DDPStrategy(find_unused_parameters=True),
72
+ )
73
+
74
+ trainer.fit(system, datamodule=datamodule)
75
+ print_only("Training finished!")
76
+ best_k = {k: v.item() for k, v in checkpoint.best_k_models.items()}
77
+ with open(os.path.join(cfg.exp.dir, cfg.exp.name, "best_k_models.json"), "w") as f:
78
+ json.dump(best_k, f, indent=0)
79
+
80
+ import wandb
81
+ if wandb.run:
82
+ print_only("Closing wandb!")
83
+ wandb.finish()
84
+
85
+ if __name__ == "__main__":
86
+
87
+ parser = argparse.ArgumentParser()
88
+ parser.add_argument(
89
+ "--conf_dir",
90
+ default="local/conf.yml",
91
+ help="Full path to save best validation model",
92
+ )
93
+
94
+ args = parser.parse_args()
95
+ cfg = OmegaConf.load(args.conf_dir)
96
+
97
+ os.makedirs(os.path.join(cfg.exp.dir, cfg.exp.name), exist_ok=True)
98
+ # 保存配置到新的文件
99
+ OmegaConf.save(cfg, os.path.join(cfg.exp.dir, cfg.exp.name, "config.yaml"))
100
+
101
+ train(cfg)
102
+
clean/audio/shiftyspeech/.env ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ WANDB_API_KEY="<wandb api-key>"
2
+ WANDB_PROJECT_NAME="SSL-AASIST"
clean/audio/shiftyspeech/LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2022 Hemlata
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
clean/audio/shiftyspeech/RawBoost.py ADDED
@@ -0,0 +1,143 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ # -*- coding: utf-8 -*-
3
+
4
+ import copy
5
+
6
+ import numpy as np
7
+ from scipy import signal
8
+
9
+ """
10
+ Hemlata Tak, Madhu Kamble, Jose Patino, Massimiliano Todisco, Nicholas Evans.
11
+ RawBoost: A Raw Data Boosting and Augmentation Method applied to Automatic Speaker Verification Anti-Spoofing.
12
+ In Proc. ICASSP 2022, pp:6382--6386.
13
+ """
14
+
15
+
16
+ def randRange(x1, x2, integer):
17
+ y = np.random.uniform(low=x1, high=x2, size=(1,))
18
+ if integer:
19
+ y = int(y)
20
+ return y
21
+
22
+
23
+ def normWav(x, always):
24
+ if always:
25
+ x = x / np.amax(abs(x))
26
+ elif np.amax(abs(x)) > 1:
27
+ x = x / np.amax(abs(x))
28
+ return x
29
+
30
+
31
+ def genNotchCoeffs(
32
+ nBands, minF, maxF, minBW, maxBW, minCoeff, maxCoeff, minG, maxG, fs
33
+ ):
34
+ b = 1
35
+ for i in range(0, nBands):
36
+ fc = randRange(minF, maxF, 0)
37
+ bw = randRange(minBW, maxBW, 0)
38
+ c = randRange(minCoeff, maxCoeff, 1)
39
+
40
+ if c / 2 == int(c / 2):
41
+ c = c + 1
42
+ f1 = fc - bw / 2
43
+ f2 = fc + bw / 2
44
+ if f1 <= 0:
45
+ f1 = 1 / 1000
46
+ if f2 >= fs / 2:
47
+ f2 = fs / 2 - 1 / 1000
48
+ b = np.convolve(
49
+ signal.firwin(c, [float(f1), float(f2)], window="hamming", fs=fs), b
50
+ )
51
+
52
+ G = randRange(minG, maxG, 0)
53
+ _, h = signal.freqz(b, 1, fs=fs)
54
+ b = pow(10, G / 20) * b / np.amax(abs(h))
55
+ return b
56
+
57
+
58
+ def filterFIR(x, b):
59
+ N = b.shape[0] + 1
60
+ xpad = np.pad(x, (0, N), "constant")
61
+ y = signal.lfilter(b, 1, xpad)
62
+ y = y[int(N / 2) : int(y.shape[0] - N / 2)]
63
+ return y
64
+
65
+
66
+ # Linear and non-linear convolutive noise
67
+ def LnL_convolutive_noise(
68
+ x,
69
+ N_f,
70
+ nBands,
71
+ minF,
72
+ maxF,
73
+ minBW,
74
+ maxBW,
75
+ minCoeff,
76
+ maxCoeff,
77
+ minG,
78
+ maxG,
79
+ minBiasLinNonLin,
80
+ maxBiasLinNonLin,
81
+ fs,
82
+ ):
83
+ y = [0] * x.shape[0]
84
+ for i in range(0, N_f):
85
+ if i == 1:
86
+ minG = minG - minBiasLinNonLin
87
+ maxG = maxG - maxBiasLinNonLin
88
+ b = genNotchCoeffs(
89
+ nBands, minF, maxF, minBW, maxBW, minCoeff, maxCoeff, minG, maxG, fs
90
+ )
91
+ y = y + filterFIR(np.power(x, (i + 1)), b)
92
+ y = y - np.mean(y)
93
+ y = normWav(y, 0)
94
+ return y
95
+
96
+
97
+ # Impulsive signal dependent noise
98
+ def ISD_additive_noise(x, P, g_sd):
99
+ beta = randRange(0, P, 0)
100
+
101
+ y = copy.deepcopy(x)
102
+ x_len = x.shape[0]
103
+ n = int(x_len * (beta / 100))
104
+ p = np.random.permutation(x_len)[:n]
105
+ f_r = np.multiply(
106
+ ((2 * np.random.rand(p.shape[0])) - 1), ((2 * np.random.rand(p.shape[0])) - 1)
107
+ )
108
+ r = g_sd * x[p] * f_r
109
+ y[p] = x[p] + r
110
+ y = normWav(y, 0)
111
+ return y
112
+
113
+
114
+ # Stationary signal independent noise
115
+
116
+
117
+ def SSI_additive_noise(
118
+ x,
119
+ SNRmin,
120
+ SNRmax,
121
+ nBands,
122
+ minF,
123
+ maxF,
124
+ minBW,
125
+ maxBW,
126
+ minCoeff,
127
+ maxCoeff,
128
+ minG,
129
+ maxG,
130
+ fs,
131
+ ):
132
+ noise = np.random.normal(0, 1, x.shape[0])
133
+ b = genNotchCoeffs(
134
+ nBands, minF, maxF, minBW, maxBW, minCoeff, maxCoeff, minG, maxG, fs
135
+ )
136
+ noise = filterFIR(noise, b)
137
+ noise = normWav(noise, 1)
138
+ SNR = randRange(SNRmin, SNRmax, 0)
139
+ noise = (
140
+ noise / np.linalg.norm(noise, 2) * np.linalg.norm(x, 2) / 10.0 ** (0.05 * SNR)
141
+ )
142
+ x = x + noise
143
+ return x
clean/audio/shiftyspeech/SOURCE.md ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Source: audio/shiftyspeech
2
+
3
+ | Field | Value |
4
+ |---|---|
5
+ | Upstream | **UNVERIFIED** -- provenance was lost when this code was vendored |
6
+ | Paper | not recorded |
7
+ | Commit SHA | **not recorded** -- the vendoring step did not preserve it |
8
+ | Mirrored on | 2026-09-15 |
9
+ | Upstream license | LICENSE |
10
+
11
+ This is a **mirror**, stripped to the files needed for inference. The full
12
+ untouched snapshot is at `archive/audio__shiftyspeech.tar.gz`.
13
+
14
+ This code is the work of its original authors and is **not** covered by the
15
+ DeepSafe project license. If you are an author and want this removed, open an
16
+ issue on https://github.com/deepsafehq/deepsafe-bench and it will be taken down
17
+ within 48 hours, no questions asked.
clean/audio/shiftyspeech/Simplified_CM_solution.py ADDED
@@ -0,0 +1,227 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ from collections import OrderedDict
3
+
4
+ import fairseq
5
+ import numpy as np
6
+ import scipy.io as sio
7
+ import torch
8
+ import torch.nn as nn
9
+ import torch.nn.functional as F
10
+ from torch import Tensor
11
+ from torch.autograd import Variable
12
+ from torch.nn.parameter import Parameter
13
+ from torch.utils import data
14
+
15
+ ___author__ = "Hemlata Tak"
16
+ __email__ = "tak@eurecom.fr"
17
+
18
+ # from losses_anti_spoofing import AMSoftmax
19
+
20
+ ############################
21
+ ## FOR fine-tuning SSL MODEL
22
+ ############################
23
+
24
+
25
+ class SSLModel(nn.Module):
26
+ def __init__(self, device):
27
+ super(SSLModel, self).__init__()
28
+
29
+ cp_path = "/change_to_path_to_pre_trained_model_XLR_300M/xlsr2_300m.pt"
30
+ model, cfg, task = fairseq.checkpoint_utils.load_model_ensemble_and_task(
31
+ [cp_path]
32
+ )
33
+ self.model = model[0]
34
+ self.device = device
35
+ self.out_dim = 1024
36
+ return
37
+
38
+ def extract_feat(self, input_data):
39
+
40
+ # put the model to GPU if it not there
41
+ if (
42
+ next(self.model.parameters()).device != input_data.device
43
+ or next(self.model.parameters()).dtype != input_data.dtype
44
+ ):
45
+ self.model.to(input_data.device, dtype=input_data.dtype)
46
+ self.model.train()
47
+
48
+ if True:
49
+ # input should be in shape (batch, length)
50
+ if input_data.ndim == 3:
51
+ input_tmp = input_data[:, :, 0]
52
+ else:
53
+ input_tmp = input_data
54
+
55
+ # [batch, length, dim]
56
+ emb = self.model(input_tmp, mask=False, features_only=True)["x"]
57
+ return emb
58
+
59
+
60
+ # ---------Graph attention simple back-end------------------------#
61
+ """
62
+ Hemlata Tak, Jee-weon Jung, Jose Patino, Madhu Kamble, Massimiliano Todisco, Nicholas Evans.
63
+ End-to-end spectro-temporal graph attention networks for speaker verification anti-spoofing and speech deepfake detection.
64
+ In Proc. Automatic Speaker Verification and Spoofing Countermeasures Challenge 2021 Interspeech 2021 satellite workshop.
65
+ """
66
+
67
+
68
+ class GraphAttentionLayer(nn.Module):
69
+ def __init__(self, in_dim, out_dim, **kwargs):
70
+ super(GraphAttentionLayer, self).__init__()
71
+
72
+ # attention map
73
+ self.att_proj = nn.Linear(in_dim, out_dim)
74
+ self.att_weight = self._init_new_params(out_dim, 1)
75
+
76
+ # project
77
+ self.proj_with_att = nn.Linear(in_dim, out_dim)
78
+ self.proj_without_att = nn.Linear(in_dim, out_dim)
79
+
80
+ # batch norm
81
+ self.bn = nn.BatchNorm1d(out_dim)
82
+
83
+ # dropout for inputs
84
+ self.input_drop = nn.Dropout(p=0.2)
85
+
86
+ self.act = nn.SELU(inplace=True)
87
+
88
+ def forward(self, x):
89
+ """
90
+ x :(#bs, #node, #dim)
91
+ """
92
+ # apply input dropout
93
+ x = self.input_drop(x)
94
+
95
+ # derive attention map
96
+ att_map = self._derive_att_map(x)
97
+
98
+ # projection
99
+ x = self._project(x, att_map)
100
+
101
+ # apply batch norm
102
+ x = self._apply_BN(x)
103
+ x = self.act(x)
104
+
105
+ return x
106
+
107
+ def _pairwise_mul_nodes(self, x):
108
+ """
109
+ Calculates pairwise multiplication of nodes.
110
+ - for attention map
111
+ x :(#bs, #node, #dim)
112
+ out_shape :(#bs, #node, #node, #dim)
113
+ """
114
+
115
+ nb_nodes = x.size(1)
116
+ x = x.unsqueeze(2).expand(-1, -1, nb_nodes, -1)
117
+ x_mirror = x.transpose(1, 2)
118
+
119
+ return x * x_mirror
120
+
121
+ def _derive_att_map(self, x):
122
+ """
123
+ x :(#bs, #node, #dim)
124
+ out_shape :(#bs, #node, #node, 1)
125
+ """
126
+ att_map = self._pairwise_mul_nodes(x)
127
+ att_map = torch.tanh(
128
+ self.att_proj(att_map)
129
+ ) # size: (#bs, #node, #node, #dim_out)
130
+ att_map = torch.matmul(att_map, self.att_weight) # size: (#bs, #node, #node, 1)
131
+ att_map = F.softmax(att_map, dim=-2)
132
+
133
+ return att_map
134
+
135
+ def _project(self, x, att_map):
136
+ x1 = self.proj_with_att(torch.matmul(att_map.squeeze(-1), x))
137
+ x2 = self.proj_without_att(x)
138
+
139
+ return x1 + x2
140
+
141
+ def _apply_BN(self, x):
142
+ org_size = x.size()
143
+ x = x.view(-1, org_size[-1])
144
+ x = self.bn(x)
145
+ x = x.view(org_size)
146
+
147
+ return x
148
+
149
+ def _init_new_params(self, *size):
150
+ out = nn.Parameter(torch.FloatTensor(*size))
151
+ nn.init.xavier_normal_(out)
152
+ return out
153
+
154
+
155
+ class GraphPool(nn.Module):
156
+ def __init__(self, k: float, in_dim: int, p):
157
+ super().__init__()
158
+ self.k = k
159
+ self.sigmoid = nn.Sigmoid()
160
+ self.proj = nn.Linear(in_dim, 1)
161
+ self.drop = nn.Dropout(p=p) if p > 0 else nn.Identity()
162
+ self.in_dim = in_dim
163
+
164
+ def forward(self, h):
165
+ Z = self.drop(h)
166
+ weights = self.proj(Z)
167
+ scores = self.sigmoid(weights)
168
+ new_h = self.top_k_graph(scores, h, self.k)
169
+
170
+ return new_h
171
+
172
+ def top_k_graph(self, scores, h, k):
173
+ """
174
+ args
175
+ =====
176
+ scores: attention-based weights (#bs, #node, 1)
177
+ h: graph data (#bs, #node, #dim)
178
+ k: ratio of remaining nodes, (float)
179
+ returns
180
+ =====
181
+ h: graph pool applied data (#bs, #node', #dim)
182
+ """
183
+ _, n_nodes, n_feat = h.size()
184
+ n_nodes = max(int(n_nodes * k), 1)
185
+ _, idx = torch.topk(scores, n_nodes, dim=1)
186
+ idx = idx.expand(-1, -1, n_feat)
187
+
188
+ h = h * scores
189
+ h = torch.gather(h, 1, idx)
190
+
191
+ return h
192
+
193
+
194
+ class Model(nn.Module):
195
+ def __init__(self, d_args, device):
196
+ super(Model, self).__init__()
197
+
198
+ # SSL model
199
+ self.device = device
200
+ self.ssl_model = SSLModel(self.device)
201
+ self.LL = nn.Linear(self.ssl_model.out_dim, 128)
202
+ self.first_bn = nn.BatchNorm1d(num_features=128)
203
+ self.selu = nn.SELU(inplace=True)
204
+
205
+ # graph module layer
206
+ self.GAT_layer = GraphAttentionLayer(128, 64)
207
+ self.proj = nn.Linear(64, 1)
208
+ self.pool = GraphPool(0.8, 64, 0.3)
209
+
210
+ # classifier head
211
+ self.proj_node = nn.Linear(53, 2)
212
+
213
+ def forward(self, x_inp, Freq_aug=False):
214
+ # SSL wav2vec 2.0 model
215
+ x_ssl_feat = self.ssl_model.extract_feat(x_inp.squeeze(-1))
216
+ x_SSL = self.LL(x_ssl_feat) # (bs,frame_number,feat_out_dim)
217
+ x_SSL = x_SSL.transpose(1, 2) # (bs,feat_out_dim,frame_number)
218
+
219
+ x = F.max_pool1d(x_SSL, (3))
220
+ x = self.first_bn(x)
221
+ x = self.selu(x)
222
+
223
+ x = self.GAT_layer(x.transpose(1, 2))
224
+ x = self.pool(x)
225
+ x = self.proj(x).flatten(1)
226
+ output = self.proj_node(x)
227
+ return output
clean/audio/shiftyspeech/data_utils.py ADDED
@@ -0,0 +1,292 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import random
3
+ from random import randrange
4
+
5
+ import librosa
6
+ import numpy as np
7
+ import torch
8
+ import torch.nn as nn
9
+ from RawBoost import (
10
+ ISD_additive_noise,
11
+ LnL_convolutive_noise,
12
+ SSI_additive_noise,
13
+ normWav,
14
+ )
15
+ from torch import Tensor
16
+ from torch.utils.data import Dataset
17
+
18
+ __author__ = "Hemlata Tak"
19
+ __email__ = "tak@eurecom.fr"
20
+
21
+
22
+ def genSpoof_list(dir_meta, is_train=False, is_eval=False):
23
+ d_meta = {}
24
+ file_list = []
25
+ with open(dir_meta, "r") as f:
26
+ l_meta = f.readlines()
27
+
28
+ if is_train:
29
+ for line in l_meta:
30
+ key, label = line.strip().split()
31
+ file_list.append(key)
32
+ d_meta[key] = 1 if label == "bonafide" else 0
33
+ return d_meta, file_list
34
+
35
+ elif is_eval:
36
+ for line in l_meta:
37
+ key, _ = line.strip().split(" ")
38
+ file_list.append(key)
39
+ return file_list
40
+ else:
41
+ for line in l_meta:
42
+ key, label = line.strip().split()
43
+ file_list.append(key)
44
+ d_meta[key] = 1 if label == "bonafide" else 0
45
+ return d_meta, file_list
46
+
47
+
48
+ def pad(x, max_len=64600):
49
+ x_len = x.shape[0]
50
+ if x_len >= max_len:
51
+ return x[:max_len]
52
+ # need to pad
53
+ num_repeats = int(max_len / x_len) + 1
54
+ padded_x = np.tile(x, (1, num_repeats))[:, :max_len][0]
55
+ return padded_x
56
+
57
+
58
+ class Dataset_ASVspoof2019_train(Dataset):
59
+ def __init__(self, args, metafile, algo):
60
+ """self.list_IDs : list of strings (each string: utt key),
61
+ self.labels: dictionary (key: utt key, value: label integer)"""
62
+
63
+ self.uttpath_labels = []
64
+ with open(metafile, "r") as f:
65
+ for line in f:
66
+ items = line.strip().split()
67
+ lb = 1 if items[-1] == "bonafide" else 0
68
+ self.uttpath_labels.append((items[0], lb))
69
+
70
+ self.algo = algo
71
+ self.args = args
72
+ self.cut = 64600 # take ~4 sec audio (64600 samples)
73
+
74
+ def __len__(self):
75
+ return len(self.uttpath_labels)
76
+
77
+ def __getitem__(self, index):
78
+ path, target = self.uttpath_labels[index]
79
+ X, fs = librosa.load(path, sr=16000)
80
+ Y = process_Rawboost_feature(X, fs, self.args, self.algo)
81
+ X_pad = pad(Y, self.cut)
82
+ x_inp = Tensor(X_pad)
83
+ return x_inp, target
84
+
85
+
86
+ class Dataset_ASVspoof2021_eval(Dataset):
87
+ def __init__(self, list_IDs):
88
+ """self.list_IDs : list of strings (each string: utt key),"""
89
+
90
+ self.list_IDs = list_IDs
91
+ self.cut = 64600 # take ~4 sec audio (64600 samples)
92
+
93
+ def __len__(self):
94
+ return len(self.list_IDs)
95
+
96
+ def __getitem__(self, index):
97
+ utt_id = self.list_IDs[index]
98
+ X, fs = librosa.load(utt_id, sr=16000)
99
+ X_pad = pad(X, self.cut)
100
+ x_inp = Tensor(X_pad)
101
+ return x_inp, utt_id
102
+
103
+
104
+ # --------------RawBoost data augmentation algorithms---------------------------##
105
+ def process_Rawboost_feature(feature, sr, args, algo):
106
+
107
+ # Data process by Convolutive noise (1st algo)
108
+ if algo == 1:
109
+
110
+ feature = LnL_convolutive_noise(
111
+ feature,
112
+ args.N_f,
113
+ args.nBands,
114
+ args.minF,
115
+ args.maxF,
116
+ args.minBW,
117
+ args.maxBW,
118
+ args.minCoeff,
119
+ args.maxCoeff,
120
+ args.minG,
121
+ args.maxG,
122
+ args.minBiasLinNonLin,
123
+ args.maxBiasLinNonLin,
124
+ sr,
125
+ )
126
+
127
+ # Data process by Impulsive noise (2nd algo)
128
+ elif algo == 2:
129
+
130
+ feature = ISD_additive_noise(feature, args.P, args.g_sd)
131
+
132
+ # Data process by coloured additive noise (3rd algo)
133
+ elif algo == 3:
134
+
135
+ feature = SSI_additive_noise(
136
+ feature,
137
+ args.SNRmin,
138
+ args.SNRmax,
139
+ args.nBands,
140
+ args.minF,
141
+ args.maxF,
142
+ args.minBW,
143
+ args.maxBW,
144
+ args.minCoeff,
145
+ args.maxCoeff,
146
+ args.minG,
147
+ args.maxG,
148
+ sr,
149
+ )
150
+
151
+ # Data process by all 3 algo. together in series (1+2+3)
152
+ elif algo == 4:
153
+
154
+ feature = LnL_convolutive_noise(
155
+ feature,
156
+ args.N_f,
157
+ args.nBands,
158
+ args.minF,
159
+ args.maxF,
160
+ args.minBW,
161
+ args.maxBW,
162
+ args.minCoeff,
163
+ args.maxCoeff,
164
+ args.minG,
165
+ args.maxG,
166
+ args.minBiasLinNonLin,
167
+ args.maxBiasLinNonLin,
168
+ sr,
169
+ )
170
+ feature = ISD_additive_noise(feature, args.P, args.g_sd)
171
+ feature = SSI_additive_noise(
172
+ feature,
173
+ args.SNRmin,
174
+ args.SNRmax,
175
+ args.nBands,
176
+ args.minF,
177
+ args.maxF,
178
+ args.minBW,
179
+ args.maxBW,
180
+ args.minCoeff,
181
+ args.maxCoeff,
182
+ args.minG,
183
+ args.maxG,
184
+ sr,
185
+ )
186
+
187
+ # Data process by 1st two algo. together in series (1+2)
188
+ elif algo == 5:
189
+
190
+ feature = LnL_convolutive_noise(
191
+ feature,
192
+ args.N_f,
193
+ args.nBands,
194
+ args.minF,
195
+ args.maxF,
196
+ args.minBW,
197
+ args.maxBW,
198
+ args.minCoeff,
199
+ args.maxCoeff,
200
+ args.minG,
201
+ args.maxG,
202
+ args.minBiasLinNonLin,
203
+ args.maxBiasLinNonLin,
204
+ sr,
205
+ )
206
+ feature = ISD_additive_noise(feature, args.P, args.g_sd)
207
+
208
+ # Data process by 1st and 3rd algo. together in series (1+3)
209
+ elif algo == 6:
210
+
211
+ feature = LnL_convolutive_noise(
212
+ feature,
213
+ args.N_f,
214
+ args.nBands,
215
+ args.minF,
216
+ args.maxF,
217
+ args.minBW,
218
+ args.maxBW,
219
+ args.minCoeff,
220
+ args.maxCoeff,
221
+ args.minG,
222
+ args.maxG,
223
+ args.minBiasLinNonLin,
224
+ args.maxBiasLinNonLin,
225
+ sr,
226
+ )
227
+ feature = SSI_additive_noise(
228
+ feature,
229
+ args.SNRmin,
230
+ args.SNRmax,
231
+ args.nBands,
232
+ args.minF,
233
+ args.maxF,
234
+ args.minBW,
235
+ args.maxBW,
236
+ args.minCoeff,
237
+ args.maxCoeff,
238
+ args.minG,
239
+ args.maxG,
240
+ sr,
241
+ )
242
+
243
+ # Data process by 2nd and 3rd algo. together in series (2+3)
244
+ elif algo == 7:
245
+
246
+ feature = ISD_additive_noise(feature, args.P, args.g_sd)
247
+ feature = SSI_additive_noise(
248
+ feature,
249
+ args.SNRmin,
250
+ args.SNRmax,
251
+ args.nBands,
252
+ args.minF,
253
+ args.maxF,
254
+ args.minBW,
255
+ args.maxBW,
256
+ args.minCoeff,
257
+ args.maxCoeff,
258
+ args.minG,
259
+ args.maxG,
260
+ sr,
261
+ )
262
+
263
+ # Data process by 1st two algo. together in Parallel (1||2)
264
+ elif algo == 8:
265
+
266
+ feature1 = LnL_convolutive_noise(
267
+ feature,
268
+ args.N_f,
269
+ args.nBands,
270
+ args.minF,
271
+ args.maxF,
272
+ args.minBW,
273
+ args.maxBW,
274
+ args.minCoeff,
275
+ args.maxCoeff,
276
+ args.minG,
277
+ args.maxG,
278
+ args.minBiasLinNonLin,
279
+ args.maxBiasLinNonLin,
280
+ sr,
281
+ )
282
+ feature2 = ISD_additive_noise(feature, args.P, args.g_sd)
283
+
284
+ feature_para = feature1 + feature2
285
+ feature = normWav(feature_para, 0) # normalized resultant waveform
286
+
287
+ # original data without Rawboost processing
288
+ else:
289
+
290
+ feature = feature
291
+
292
+ return feature
clean/audio/shiftyspeech/model.py ADDED
@@ -0,0 +1,603 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import random
2
+ from typing import Union
3
+
4
+ import fairseq
5
+ import numpy as np
6
+ import torch
7
+ import torch.nn as nn
8
+ import torch.nn.functional as F
9
+ from torch import Tensor
10
+
11
+ ___author__ = "Hemlata Tak"
12
+ __email__ = "tak@eurecom.fr"
13
+
14
+ ############################
15
+ ## FOR fine-tuned SSL MODEL
16
+ ############################
17
+
18
+
19
+ class SSLModel(nn.Module):
20
+ def __init__(self, device):
21
+ super(SSLModel, self).__init__()
22
+
23
+ cp_path = "models/xlsr2_300m.pt" # Change the pre-trained XLSR model path.
24
+ model, cfg, task = fairseq.checkpoint_utils.load_model_ensemble_and_task(
25
+ [cp_path]
26
+ )
27
+ self.model = model[0]
28
+ self.device = device
29
+ self.out_dim = 1024
30
+ return
31
+
32
+ def extract_feat(self, input_data):
33
+
34
+ # put the model to GPU if it not there
35
+ if (
36
+ next(self.model.parameters()).device != input_data.device
37
+ or next(self.model.parameters()).dtype != input_data.dtype
38
+ ):
39
+ self.model.to(input_data.device, dtype=input_data.dtype)
40
+ self.model.train()
41
+
42
+ if True:
43
+ # input should be in shape (batch, length)
44
+ if input_data.ndim == 3:
45
+ input_tmp = input_data[:, :, 0]
46
+ else:
47
+ input_tmp = input_data
48
+
49
+ # [batch, length, dim]
50
+ emb = self.model(input_tmp, mask=False, features_only=True)["x"]
51
+ return emb
52
+
53
+
54
+ # ---------AASIST back-end------------------------#
55
+ """ Jee-weon Jung, Hee-Soo Heo, Hemlata Tak, Hye-jin Shim, Joon Son Chung, Bong-Jin Lee, Ha-Jin Yu and Nicholas Evans.
56
+ AASIST: Audio Anti-Spoofing Using Integrated Spectro-Temporal Graph Attention Networks.
57
+ In Proc. ICASSP 2022, pp: 6367--6371."""
58
+
59
+
60
+ class GraphAttentionLayer(nn.Module):
61
+ def __init__(self, in_dim, out_dim, **kwargs):
62
+ super().__init__()
63
+
64
+ # attention map
65
+ self.att_proj = nn.Linear(in_dim, out_dim)
66
+ self.att_weight = self._init_new_params(out_dim, 1)
67
+
68
+ # project
69
+ self.proj_with_att = nn.Linear(in_dim, out_dim)
70
+ self.proj_without_att = nn.Linear(in_dim, out_dim)
71
+
72
+ # batch norm
73
+ self.bn = nn.BatchNorm1d(out_dim)
74
+
75
+ # dropout for inputs
76
+ self.input_drop = nn.Dropout(p=0.2)
77
+
78
+ # activate
79
+ self.act = nn.SELU(inplace=True)
80
+
81
+ # temperature
82
+ self.temp = 1.0
83
+ if "temperature" in kwargs:
84
+ self.temp = kwargs["temperature"]
85
+
86
+ def forward(self, x):
87
+ """
88
+ x :(#bs, #node, #dim)
89
+ """
90
+ # apply input dropout
91
+ x = self.input_drop(x)
92
+
93
+ # derive attention map
94
+ att_map = self._derive_att_map(x)
95
+
96
+ # projection
97
+ x = self._project(x, att_map)
98
+
99
+ # apply batch norm
100
+ x = self._apply_BN(x)
101
+ x = self.act(x)
102
+ return x
103
+
104
+ def _pairwise_mul_nodes(self, x):
105
+ """
106
+ Calculates pairwise multiplication of nodes.
107
+ - for attention map
108
+ x :(#bs, #node, #dim)
109
+ out_shape :(#bs, #node, #node, #dim)
110
+ """
111
+
112
+ nb_nodes = x.size(1)
113
+ x = x.unsqueeze(2).expand(-1, -1, nb_nodes, -1)
114
+ x_mirror = x.transpose(1, 2)
115
+
116
+ return x * x_mirror
117
+
118
+ def _derive_att_map(self, x):
119
+ """
120
+ x :(#bs, #node, #dim)
121
+ out_shape :(#bs, #node, #node, 1)
122
+ """
123
+ att_map = self._pairwise_mul_nodes(x)
124
+ # size: (#bs, #node, #node, #dim_out)
125
+ att_map = torch.tanh(self.att_proj(att_map))
126
+ # size: (#bs, #node, #node, 1)
127
+ att_map = torch.matmul(att_map, self.att_weight)
128
+
129
+ # apply temperature
130
+ att_map = att_map / self.temp
131
+
132
+ att_map = F.softmax(att_map, dim=-2)
133
+
134
+ return att_map
135
+
136
+ def _project(self, x, att_map):
137
+ x1 = self.proj_with_att(torch.matmul(att_map.squeeze(-1), x))
138
+ x2 = self.proj_without_att(x)
139
+
140
+ return x1 + x2
141
+
142
+ def _apply_BN(self, x):
143
+ org_size = x.size()
144
+ x = x.view(-1, org_size[-1])
145
+ x = self.bn(x)
146
+ x = x.view(org_size)
147
+
148
+ return x
149
+
150
+ def _init_new_params(self, *size):
151
+ out = nn.Parameter(torch.FloatTensor(*size))
152
+ nn.init.xavier_normal_(out)
153
+ return out
154
+
155
+
156
+ class HtrgGraphAttentionLayer(nn.Module):
157
+ def __init__(self, in_dim, out_dim, **kwargs):
158
+ super().__init__()
159
+
160
+ self.proj_type1 = nn.Linear(in_dim, in_dim)
161
+ self.proj_type2 = nn.Linear(in_dim, in_dim)
162
+
163
+ # attention map
164
+ self.att_proj = nn.Linear(in_dim, out_dim)
165
+ self.att_projM = nn.Linear(in_dim, out_dim)
166
+
167
+ self.att_weight11 = self._init_new_params(out_dim, 1)
168
+ self.att_weight22 = self._init_new_params(out_dim, 1)
169
+ self.att_weight12 = self._init_new_params(out_dim, 1)
170
+ self.att_weightM = self._init_new_params(out_dim, 1)
171
+
172
+ # project
173
+ self.proj_with_att = nn.Linear(in_dim, out_dim)
174
+ self.proj_without_att = nn.Linear(in_dim, out_dim)
175
+
176
+ self.proj_with_attM = nn.Linear(in_dim, out_dim)
177
+ self.proj_without_attM = nn.Linear(in_dim, out_dim)
178
+
179
+ # batch norm
180
+ self.bn = nn.BatchNorm1d(out_dim)
181
+
182
+ # dropout for inputs
183
+ self.input_drop = nn.Dropout(p=0.2)
184
+
185
+ # activate
186
+ self.act = nn.SELU(inplace=True)
187
+
188
+ # temperature
189
+ self.temp = 1.0
190
+ if "temperature" in kwargs:
191
+ self.temp = kwargs["temperature"]
192
+
193
+ def forward(self, x1, x2, master=None):
194
+ """
195
+ x1 :(#bs, #node, #dim)
196
+ x2 :(#bs, #node, #dim)
197
+ """
198
+
199
+ num_type1 = x1.size(1)
200
+ num_type2 = x2.size(1)
201
+
202
+ x1 = self.proj_type1(x1)
203
+
204
+ x2 = self.proj_type2(x2)
205
+
206
+ x = torch.cat([x1, x2], dim=1)
207
+
208
+ if master is None:
209
+ master = torch.mean(x, dim=1, keepdim=True)
210
+
211
+ # apply input dropout
212
+ x = self.input_drop(x)
213
+
214
+ # derive attention map
215
+ att_map = self._derive_att_map(x, num_type1, num_type2)
216
+
217
+ # directional edge for master node
218
+ master = self._update_master(x, master)
219
+
220
+ # projection
221
+ x = self._project(x, att_map)
222
+
223
+ # apply batch norm
224
+ x = self._apply_BN(x)
225
+ x = self.act(x)
226
+
227
+ x1 = x.narrow(1, 0, num_type1)
228
+
229
+ x2 = x.narrow(1, num_type1, num_type2)
230
+
231
+ return x1, x2, master
232
+
233
+ def _update_master(self, x, master):
234
+
235
+ att_map = self._derive_att_map_master(x, master)
236
+ master = self._project_master(x, master, att_map)
237
+
238
+ return master
239
+
240
+ def _pairwise_mul_nodes(self, x):
241
+ """
242
+ Calculates pairwise multiplication of nodes.
243
+ - for attention map
244
+ x :(#bs, #node, #dim)
245
+ out_shape :(#bs, #node, #node, #dim)
246
+ """
247
+
248
+ nb_nodes = x.size(1)
249
+ x = x.unsqueeze(2).expand(-1, -1, nb_nodes, -1)
250
+ x_mirror = x.transpose(1, 2)
251
+
252
+ return x * x_mirror
253
+
254
+ def _derive_att_map_master(self, x, master):
255
+ """
256
+ x :(#bs, #node, #dim)
257
+ out_shape :(#bs, #node, #node, 1)
258
+ """
259
+ att_map = x * master
260
+ att_map = torch.tanh(self.att_projM(att_map))
261
+
262
+ att_map = torch.matmul(att_map, self.att_weightM)
263
+
264
+ # apply temperature
265
+ att_map = att_map / self.temp
266
+
267
+ att_map = F.softmax(att_map, dim=-2)
268
+
269
+ return att_map
270
+
271
+ def _derive_att_map(self, x, num_type1, num_type2):
272
+ """
273
+ x :(#bs, #node, #dim)
274
+ out_shape :(#bs, #node, #node, 1)
275
+ """
276
+ att_map = self._pairwise_mul_nodes(x)
277
+ # size: (#bs, #node, #node, #dim_out)
278
+ att_map = torch.tanh(self.att_proj(att_map))
279
+ # size: (#bs, #node, #node, 1)
280
+
281
+ att_board = torch.zeros_like(att_map[:, :, :, 0]).unsqueeze(-1)
282
+
283
+ att_board[:, :num_type1, :num_type1, :] = torch.matmul(
284
+ att_map[:, :num_type1, :num_type1, :], self.att_weight11
285
+ )
286
+ att_board[:, num_type1:, num_type1:, :] = torch.matmul(
287
+ att_map[:, num_type1:, num_type1:, :], self.att_weight22
288
+ )
289
+ att_board[:, :num_type1, num_type1:, :] = torch.matmul(
290
+ att_map[:, :num_type1, num_type1:, :], self.att_weight12
291
+ )
292
+ att_board[:, num_type1:, :num_type1, :] = torch.matmul(
293
+ att_map[:, num_type1:, :num_type1, :], self.att_weight12
294
+ )
295
+
296
+ att_map = att_board
297
+
298
+ # apply temperature
299
+ att_map = att_map / self.temp
300
+
301
+ att_map = F.softmax(att_map, dim=-2)
302
+
303
+ return att_map
304
+
305
+ def _project(self, x, att_map):
306
+ x1 = self.proj_with_att(torch.matmul(att_map.squeeze(-1), x))
307
+ x2 = self.proj_without_att(x)
308
+
309
+ return x1 + x2
310
+
311
+ def _project_master(self, x, master, att_map):
312
+
313
+ x1 = self.proj_with_attM(torch.matmul(att_map.squeeze(-1).unsqueeze(1), x))
314
+ x2 = self.proj_without_attM(master)
315
+
316
+ return x1 + x2
317
+
318
+ def _apply_BN(self, x):
319
+ org_size = x.size()
320
+ x = x.view(-1, org_size[-1])
321
+ x = self.bn(x)
322
+ x = x.view(org_size)
323
+
324
+ return x
325
+
326
+ def _init_new_params(self, *size):
327
+ out = nn.Parameter(torch.FloatTensor(*size))
328
+ nn.init.xavier_normal_(out)
329
+ return out
330
+
331
+
332
+ class GraphPool(nn.Module):
333
+ def __init__(self, k: float, in_dim: int, p: Union[float, int]):
334
+ super().__init__()
335
+ self.k = k
336
+ self.sigmoid = nn.Sigmoid()
337
+ self.proj = nn.Linear(in_dim, 1)
338
+ self.drop = nn.Dropout(p=p) if p > 0 else nn.Identity()
339
+ self.in_dim = in_dim
340
+
341
+ def forward(self, h):
342
+ Z = self.drop(h)
343
+ weights = self.proj(Z)
344
+ scores = self.sigmoid(weights)
345
+ new_h = self.top_k_graph(scores, h, self.k)
346
+
347
+ return new_h
348
+
349
+ def top_k_graph(self, scores, h, k):
350
+ """
351
+ args
352
+ =====
353
+ scores: attention-based weights (#bs, #node, 1)
354
+ h: graph data (#bs, #node, #dim)
355
+ k: ratio of remaining nodes, (float)
356
+ returns
357
+ =====
358
+ h: graph pool applied data (#bs, #node', #dim)
359
+ """
360
+ _, n_nodes, n_feat = h.size()
361
+ n_nodes = max(int(n_nodes * k), 1)
362
+ _, idx = torch.topk(scores, n_nodes, dim=1)
363
+ idx = idx.expand(-1, -1, n_feat)
364
+
365
+ h = h * scores
366
+ h = torch.gather(h, 1, idx)
367
+
368
+ return h
369
+
370
+
371
+ class Residual_block(nn.Module):
372
+ def __init__(self, nb_filts, first=False):
373
+ super().__init__()
374
+ self.first = first
375
+
376
+ if not self.first:
377
+ self.bn1 = nn.BatchNorm2d(num_features=nb_filts[0])
378
+ self.conv1 = nn.Conv2d(
379
+ in_channels=nb_filts[0],
380
+ out_channels=nb_filts[1],
381
+ kernel_size=(2, 3),
382
+ padding=(1, 1),
383
+ stride=1,
384
+ )
385
+ self.selu = nn.SELU(inplace=True)
386
+
387
+ self.bn2 = nn.BatchNorm2d(num_features=nb_filts[1])
388
+ self.conv2 = nn.Conv2d(
389
+ in_channels=nb_filts[1],
390
+ out_channels=nb_filts[1],
391
+ kernel_size=(2, 3),
392
+ padding=(0, 1),
393
+ stride=1,
394
+ )
395
+
396
+ if nb_filts[0] != nb_filts[1]:
397
+ self.downsample = True
398
+ self.conv_downsample = nn.Conv2d(
399
+ in_channels=nb_filts[0],
400
+ out_channels=nb_filts[1],
401
+ padding=(0, 1),
402
+ kernel_size=(1, 3),
403
+ stride=1,
404
+ )
405
+
406
+ else:
407
+ self.downsample = False
408
+
409
+ def forward(self, x):
410
+ identity = x
411
+ if not self.first:
412
+ out = self.bn1(x)
413
+ out = self.selu(out)
414
+ else:
415
+ out = x
416
+
417
+ out = self.conv1(x)
418
+
419
+ out = self.bn2(out)
420
+ out = self.selu(out)
421
+
422
+ out = self.conv2(out)
423
+
424
+ if self.downsample:
425
+ identity = self.conv_downsample(identity)
426
+
427
+ out += identity
428
+
429
+ return out
430
+
431
+
432
+ class Model(nn.Module):
433
+ def __init__(self, args, device):
434
+ super().__init__()
435
+ self.device = device
436
+
437
+ # AASIST parameters
438
+ filts = [128, [1, 32], [32, 32], [32, 64], [64, 64]]
439
+ gat_dims = [64, 32]
440
+ pool_ratios = [0.5, 0.5, 0.5, 0.5]
441
+ temperatures = [2.0, 2.0, 100.0, 100.0]
442
+
443
+ ####
444
+ # create network wav2vec 2.0
445
+ ####
446
+ self.ssl_model = SSLModel(self.device)
447
+ self.LL = nn.Linear(self.ssl_model.out_dim, 128)
448
+
449
+ self.first_bn = nn.BatchNorm2d(num_features=1)
450
+ self.first_bn1 = nn.BatchNorm2d(num_features=64)
451
+ self.drop = nn.Dropout(0.5, inplace=True)
452
+ self.drop_way = nn.Dropout(0.2, inplace=True)
453
+ self.selu = nn.SELU(inplace=True)
454
+
455
+ # RawNet2 encoder
456
+ self.encoder = nn.Sequential(
457
+ nn.Sequential(Residual_block(nb_filts=filts[1], first=True)),
458
+ nn.Sequential(Residual_block(nb_filts=filts[2])),
459
+ nn.Sequential(Residual_block(nb_filts=filts[3])),
460
+ nn.Sequential(Residual_block(nb_filts=filts[4])),
461
+ nn.Sequential(Residual_block(nb_filts=filts[4])),
462
+ nn.Sequential(Residual_block(nb_filts=filts[4])),
463
+ )
464
+
465
+ self.attention = nn.Sequential(
466
+ nn.Conv2d(64, 128, kernel_size=(1, 1)),
467
+ nn.SELU(inplace=True),
468
+ nn.BatchNorm2d(128),
469
+ nn.Conv2d(128, 64, kernel_size=(1, 1)),
470
+ )
471
+ # position encoding
472
+ self.pos_S = nn.Parameter(torch.randn(1, 42, filts[-1][-1]))
473
+
474
+ self.master1 = nn.Parameter(torch.randn(1, 1, gat_dims[0]))
475
+ self.master2 = nn.Parameter(torch.randn(1, 1, gat_dims[0]))
476
+
477
+ # Graph module
478
+ self.GAT_layer_S = GraphAttentionLayer(
479
+ filts[-1][-1], gat_dims[0], temperature=temperatures[0]
480
+ )
481
+ self.GAT_layer_T = GraphAttentionLayer(
482
+ filts[-1][-1], gat_dims[0], temperature=temperatures[1]
483
+ )
484
+ # HS-GAL layer
485
+ self.HtrgGAT_layer_ST11 = HtrgGraphAttentionLayer(
486
+ gat_dims[0], gat_dims[1], temperature=temperatures[2]
487
+ )
488
+ self.HtrgGAT_layer_ST12 = HtrgGraphAttentionLayer(
489
+ gat_dims[1], gat_dims[1], temperature=temperatures[2]
490
+ )
491
+ self.HtrgGAT_layer_ST21 = HtrgGraphAttentionLayer(
492
+ gat_dims[0], gat_dims[1], temperature=temperatures[2]
493
+ )
494
+ self.HtrgGAT_layer_ST22 = HtrgGraphAttentionLayer(
495
+ gat_dims[1], gat_dims[1], temperature=temperatures[2]
496
+ )
497
+
498
+ # Graph pooling layers
499
+ self.pool_S = GraphPool(pool_ratios[0], gat_dims[0], 0.3)
500
+ self.pool_T = GraphPool(pool_ratios[1], gat_dims[0], 0.3)
501
+ self.pool_hS1 = GraphPool(pool_ratios[2], gat_dims[1], 0.3)
502
+ self.pool_hT1 = GraphPool(pool_ratios[2], gat_dims[1], 0.3)
503
+
504
+ self.pool_hS2 = GraphPool(pool_ratios[2], gat_dims[1], 0.3)
505
+ self.pool_hT2 = GraphPool(pool_ratios[2], gat_dims[1], 0.3)
506
+
507
+ self.out_layer = nn.Linear(5 * gat_dims[1], 2)
508
+
509
+ def forward(self, x):
510
+ # -------pre-trained Wav2vec model fine tunning ------------------------##
511
+ x_ssl_feat = self.ssl_model.extract_feat(x.squeeze(-1))
512
+ x = self.LL(x_ssl_feat) # (bs,frame_number,feat_out_dim)
513
+
514
+ # post-processing on front-end features
515
+ x = x.transpose(1, 2) # (bs,feat_out_dim,frame_number)
516
+ x = x.unsqueeze(dim=1) # add channel
517
+ x = F.max_pool2d(x, (3, 3))
518
+ x = self.first_bn(x)
519
+ x = self.selu(x)
520
+
521
+ # RawNet2-based encoder
522
+ x = self.encoder(x)
523
+ x = self.first_bn1(x)
524
+ x = self.selu(x)
525
+
526
+ w = self.attention(x)
527
+
528
+ # ------------SA for spectral feature-------------#
529
+ w1 = F.softmax(w, dim=-1)
530
+ m = torch.sum(x * w1, dim=-1)
531
+ e_S = m.transpose(1, 2) + self.pos_S
532
+
533
+ # graph module layer
534
+ gat_S = self.GAT_layer_S(e_S)
535
+ out_S = self.pool_S(gat_S) # (#bs, #node, #dim)
536
+
537
+ # ------------SA for temporal feature-------------#
538
+ w2 = F.softmax(w, dim=-2)
539
+ m1 = torch.sum(x * w2, dim=-2)
540
+
541
+ e_T = m1.transpose(1, 2)
542
+
543
+ # graph module layer
544
+ gat_T = self.GAT_layer_T(e_T)
545
+ out_T = self.pool_T(gat_T)
546
+
547
+ # learnable master node
548
+ master1 = self.master1.expand(x.size(0), -1, -1)
549
+ master2 = self.master2.expand(x.size(0), -1, -1)
550
+
551
+ # inference 1
552
+ out_T1, out_S1, master1 = self.HtrgGAT_layer_ST11(
553
+ out_T, out_S, master=self.master1
554
+ )
555
+
556
+ out_S1 = self.pool_hS1(out_S1)
557
+ out_T1 = self.pool_hT1(out_T1)
558
+
559
+ out_T_aug, out_S_aug, master_aug = self.HtrgGAT_layer_ST12(
560
+ out_T1, out_S1, master=master1
561
+ )
562
+ out_T1 = out_T1 + out_T_aug
563
+ out_S1 = out_S1 + out_S_aug
564
+ master1 = master1 + master_aug
565
+
566
+ # inference 2
567
+ out_T2, out_S2, master2 = self.HtrgGAT_layer_ST21(
568
+ out_T, out_S, master=self.master2
569
+ )
570
+ out_S2 = self.pool_hS2(out_S2)
571
+ out_T2 = self.pool_hT2(out_T2)
572
+
573
+ out_T_aug, out_S_aug, master_aug = self.HtrgGAT_layer_ST22(
574
+ out_T2, out_S2, master=master2
575
+ )
576
+ out_T2 = out_T2 + out_T_aug
577
+ out_S2 = out_S2 + out_S_aug
578
+ master2 = master2 + master_aug
579
+
580
+ out_T1 = self.drop_way(out_T1)
581
+ out_T2 = self.drop_way(out_T2)
582
+ out_S1 = self.drop_way(out_S1)
583
+ out_S2 = self.drop_way(out_S2)
584
+ master1 = self.drop_way(master1)
585
+ master2 = self.drop_way(master2)
586
+
587
+ out_T = torch.max(out_T1, out_T2)
588
+ out_S = torch.max(out_S1, out_S2)
589
+ master = torch.max(master1, master2)
590
+
591
+ # Readout operation
592
+ T_max, _ = torch.max(torch.abs(out_T), dim=1)
593
+ T_avg = torch.mean(out_T, dim=1)
594
+
595
+ S_max, _ = torch.max(torch.abs(out_S), dim=1)
596
+ S_avg = torch.mean(out_S, dim=1)
597
+
598
+ last_hidden = torch.cat([T_max, T_avg, S_max, S_avg, master.squeeze(1)], dim=1)
599
+
600
+ last_hidden = self.drop(last_hidden)
601
+ output = self.out_layer(last_hidden)
602
+
603
+ return output
clean/audio/shiftyspeech/startup_config.py ADDED
@@ -0,0 +1,60 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """
3
+ startup_config
4
+
5
+ Startup configuration utilities
6
+
7
+ """
8
+
9
+ from __future__ import absolute_import
10
+
11
+ import importlib
12
+ import os
13
+ import random
14
+ import sys
15
+
16
+ import numpy as np
17
+ import torch
18
+
19
+ __author__ = "Xin Wang"
20
+ __email__ = "wangxin@nii.ac.jp"
21
+ __copyright__ = "Copyright 2020, Xin Wang"
22
+
23
+
24
+ def set_random_seed(random_seed, args=None):
25
+ """set_random_seed(random_seed, args=None)
26
+
27
+ Set the random_seed for numpy, python, and cudnn
28
+
29
+ input
30
+ -----
31
+ random_seed: integer random seed
32
+ args: argue parser
33
+ """
34
+
35
+ # initialization
36
+ torch.manual_seed(random_seed)
37
+ random.seed(random_seed)
38
+ np.random.seed(random_seed)
39
+ os.environ["PYTHONHASHSEED"] = str(random_seed)
40
+
41
+ # For torch.backends.cudnn.deterministic
42
+ # Note: this default configuration may result in RuntimeError
43
+ # see https://pytorch.org/docs/stable/notes/randomness.html
44
+ if args is None:
45
+ cudnn_deterministic = True
46
+ cudnn_benchmark = False
47
+ else:
48
+ cudnn_deterministic = args.cudnn_deterministic_toggle
49
+ cudnn_benchmark = args.cudnn_benchmark_toggle
50
+
51
+ if not cudnn_deterministic:
52
+ print("cudnn_deterministic set to False")
53
+ if cudnn_benchmark:
54
+ print("cudnn_benchmark set to True")
55
+
56
+ if torch.cuda.is_available():
57
+ torch.cuda.manual_seed_all(random_seed)
58
+ torch.backends.cudnn.deterministic = cudnn_deterministic
59
+ torch.backends.cudnn.benchmark = cudnn_benchmark
60
+ return
clean/audio/shiftyspeech/train.py ADDED
@@ -0,0 +1,446 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import os
3
+ import sys
4
+
5
+ import librosa
6
+ import numpy as np
7
+ import torch
8
+ import wandb
9
+ import yaml
10
+ from data_utils import (
11
+ Dataset_ASVspoof2019_train,
12
+ Dataset_ASVspoof2021_eval,
13
+ genSpoof_list,
14
+ pad,
15
+ process_Rawboost_feature,
16
+ )
17
+ from dotenv import load_dotenv
18
+ from model import Model
19
+ from sklearn.metrics import roc_auc_score
20
+ from startup_config import set_random_seed
21
+ from tensorboardX import SummaryWriter
22
+ from torch import Tensor, nn
23
+ from torch.utils.data import DataLoader
24
+ from tqdm import tqdm
25
+
26
+ __author__ = "Hemlata Tak"
27
+ __email__ = "tak@eurecom.fr"
28
+
29
+
30
+ def compute_det_curve(target_scores, nontarget_scores):
31
+
32
+ n_scores = target_scores.size + nontarget_scores.size
33
+ all_scores = np.concatenate((target_scores, nontarget_scores))
34
+ labels = np.concatenate(
35
+ (np.ones(target_scores.size), np.zeros(nontarget_scores.size))
36
+ )
37
+
38
+ indices = np.argsort(all_scores, kind="mergesort")
39
+ labels = labels[indices]
40
+ tar_trial_sums = np.cumsum(labels)
41
+ nontarget_trial_sums = nontarget_scores.size - (
42
+ np.arange(1, n_scores + 1) - tar_trial_sums
43
+ )
44
+
45
+ frr = np.concatenate((np.atleast_1d(0), tar_trial_sums / target_scores.size))
46
+ far = np.concatenate(
47
+ (np.atleast_1d(1), nontarget_trial_sums / nontarget_scores.size)
48
+ )
49
+ # Thresholds are the sorted scores
50
+ thresholds = np.concatenate(
51
+ (np.atleast_1d(all_scores[indices[0]] - 0.001), all_scores[indices])
52
+ )
53
+
54
+ return frr, far, thresholds
55
+
56
+
57
+ def compute_eer(target_scores, nontarget_scores):
58
+ """Returns equal error rate (EER) and the corresponding threshold."""
59
+ frr, far, thresholds = compute_det_curve(target_scores, nontarget_scores)
60
+ abs_diffs = np.abs(frr - far)
61
+ min_index = np.argmin(abs_diffs)
62
+ eer = np.mean((frr[min_index], far[min_index]))
63
+ return eer, thresholds[min_index], frr, far
64
+
65
+
66
+ def calculate_tDCF_EER(cm_scores_file, output_file, printout=True):
67
+ # Load CM scores
68
+ cm_data = np.genfromtxt(cm_scores_file, dtype=str)
69
+ cm_utt_id = cm_data[:, 0]
70
+ cm_keys = cm_data[:, 1]
71
+ cm_scores = cm_data[:, 2].astype(float)
72
+ # Extract bona fide (real human) and spoof scores from the CM scores
73
+ bona_cm = cm_scores[cm_keys == "bonafide"]
74
+ spoof_cm = cm_scores[cm_keys == "spoof"]
75
+ all_scores = np.concatenate([bona_cm, spoof_cm])
76
+ all_true_labels = np.concatenate([np.ones_like(bona_cm), np.zeros_like(spoof_cm)])
77
+
78
+ auc = roc_auc_score(all_true_labels, all_scores, max_fpr=0.05)
79
+ eer_cm, eer_threshold, frr, far = compute_eer(bona_cm, spoof_cm)
80
+
81
+ if printout:
82
+ with open(output_file, "w") as f_res:
83
+ f_res.write("\nCM SYSTEM\n")
84
+ f_res.write(
85
+ "\tEER\t\t= {:8.9f} % "
86
+ "(Equal error rate for countermeasure)\n".format(eer_cm * 100)
87
+ )
88
+ f_res.write("\t pAUC with max fpr - 0.05 is :{}".format(auc))
89
+
90
+
91
+ def evaluate_accuracy(dev_loader, model, device, args):
92
+ val_loss = 0.0
93
+ num_total = 0.0
94
+ algo = args.algo
95
+ cut = 64600
96
+ model.eval()
97
+
98
+ weight = torch.FloatTensor([0.1, 0.9]).to(device)
99
+ criterion = nn.CrossEntropyLoss(weight=weight)
100
+ progress_bar = tqdm(dev_loader, desc=f"Epoch {epoch+1}/{args.num_epochs}")
101
+ for current_step, (batch_pths, batch_y) in enumerate(progress_bar):
102
+ batch_x = batch_pths
103
+ batch_size = batch_x.size(0)
104
+ num_total += batch_size
105
+ batch_x = batch_x.to(device)
106
+ batch_y = batch_y.view(-1).type(torch.int64).to(device)
107
+ batch_out = model(batch_x)
108
+
109
+ batch_loss = criterion(batch_out, batch_y)
110
+ val_loss += batch_loss.item() * batch_size
111
+
112
+ val_loss /= num_total
113
+
114
+ return val_loss
115
+
116
+
117
+ def produce_evaluation_file(dataset, model, device, save_path, trial_path):
118
+ data_loader = DataLoader(dataset, batch_size=10, shuffle=False, drop_last=False)
119
+ num_correct = 0.0
120
+ num_total = 0.0
121
+ model.eval()
122
+ with open(trial_path, "r") as f_trl:
123
+ trial_lines = f_trl.readlines()
124
+
125
+ fname_list = []
126
+ score_list = []
127
+
128
+ for batch_x, utt_id in data_loader:
129
+
130
+ batch_size = batch_x.size(0)
131
+ batch_x = batch_x.to(device)
132
+
133
+ batch_out = model(batch_x)
134
+
135
+ batch_score = (batch_out[:, 1]).data.cpu().numpy().ravel()
136
+ # add outputs
137
+ fname_list.extend(utt_id)
138
+ score_list.extend(batch_score.tolist())
139
+ assert len(trial_lines) == len(fname_list) == len(score_list)
140
+
141
+ with open(save_path, "a+") as fh:
142
+ for fname, cm, trl in zip(fname_list, score_list, trial_lines):
143
+ utt_id, key = trl.strip().split(" ")
144
+ assert fname == utt_id
145
+ fh.write("{} {} {}\n".format(fname, key, cm))
146
+ fh.close()
147
+ print("Scores saved to {}".format(save_path))
148
+
149
+
150
+ def train_epoch(train_loader, model, lr, optim, device, args):
151
+ running_loss = 0
152
+
153
+ num_total = 0.0
154
+ algo = args.algo
155
+ model.train()
156
+ cut = 64600
157
+ # set objective (Loss) functions
158
+ weight = torch.FloatTensor([0.1, 0.9]).to(device)
159
+ criterion = nn.CrossEntropyLoss(weight=weight)
160
+ progress_bar = tqdm(train_loader, desc=f"Epoch {epoch+1}/{args.num_epochs}")
161
+ for current_step, (batch_pths, batch_y) in enumerate(progress_bar):
162
+ batch_x = batch_pths
163
+ batch_size = batch_x.size(0)
164
+ num_total += batch_size
165
+
166
+ batch_x = batch_x.to(device)
167
+ batch_y = batch_y.view(-1).type(torch.int64).to(device)
168
+ batch_out = model(batch_x)
169
+
170
+ batch_loss = criterion(batch_out, batch_y)
171
+
172
+ running_loss += batch_loss.item() * batch_size
173
+
174
+ optimizer.zero_grad()
175
+ batch_loss.backward()
176
+ optimizer.step()
177
+
178
+ running_loss /= num_total
179
+
180
+ return running_loss
181
+
182
+
183
+ if __name__ == "__main__":
184
+ parser = argparse.ArgumentParser(description="SSL-AASIST baseline system")
185
+
186
+ # Hyperparameters
187
+ parser.add_argument("--batch_size", type=int, default=64)
188
+ parser.add_argument("--num_epochs", type=int, default=100)
189
+ parser.add_argument("--lr", type=float, default=0.000001)
190
+ parser.add_argument("--weight_decay", type=float, default=0.0001)
191
+ parser.add_argument("--model_name", type=str, default="SSL-AASIST")
192
+ parser.add_argument("--loss", type=str, default="weighted_CCE")
193
+ parser.add_argument("--trn_list_path", default=None, help="path to train file")
194
+ parser.add_argument("--dev_list_path", default=None, help="path to validation file")
195
+ parser.add_argument("--test_list_path", default=None, help="path to test file")
196
+ parser.add_argument(
197
+ "--test_score_dir", default=None, help="path to save test scores"
198
+ )
199
+ # model
200
+ parser.add_argument(
201
+ "--seed", type=int, default=1234, help="random seed (default: 1234)"
202
+ )
203
+ parser.add_argument("--save_path", type=str, default=".", help="Model save path")
204
+ parser.add_argument("--model_path", type=str, default=None, help="Model checkpoint")
205
+ parser.add_argument(
206
+ "--comment", type=str, default=None, help="Comment to describe the saved model"
207
+ )
208
+ # Auxiliary arguments
209
+
210
+ parser.add_argument("--eval", action="store_true", default=False, help="eval mode")
211
+ parser.add_argument("--eval_part", type=int, default=0)
212
+ # backend options
213
+ parser.add_argument(
214
+ "--cudnn-deterministic-toggle",
215
+ action="store_false",
216
+ default=True,
217
+ help="use cudnn-deterministic? (default true)",
218
+ )
219
+
220
+ parser.add_argument(
221
+ "--cudnn-benchmark-toggle",
222
+ action="store_true",
223
+ default=False,
224
+ help="use cudnn-benchmark? (default false)",
225
+ )
226
+
227
+ ##===================================================Rawboost data augmentation ======================================================================#
228
+
229
+ parser.add_argument(
230
+ "--algo",
231
+ type=int,
232
+ default=5,
233
+ help="Rawboost algos discriptions. 0: No augmentation 1: LnL_convolutive_noise, 2: ISD_additive_noise, 3: SSI_additive_noise, 4: series algo (1+2+3), \
234
+ 5: series algo (1+2), 6: series algo (1+3), 7: series algo(2+3), 8: parallel algo(1,2) .[default=0]",
235
+ )
236
+
237
+ # LnL_convolutive_noise parameters
238
+ parser.add_argument(
239
+ "--nBands",
240
+ type=int,
241
+ default=5,
242
+ help="number of notch filters.The higher the number of bands, the more aggresive the distortions is.[default=5]",
243
+ )
244
+ parser.add_argument(
245
+ "--minF",
246
+ type=int,
247
+ default=20,
248
+ help="minimum centre frequency [Hz] of notch filter.[default=20] ",
249
+ )
250
+ parser.add_argument(
251
+ "--maxF",
252
+ type=int,
253
+ default=8000,
254
+ help="maximum centre frequency [Hz] (<sr/2) of notch filter.[default=8000]",
255
+ )
256
+ parser.add_argument(
257
+ "--minBW",
258
+ type=int,
259
+ default=100,
260
+ help="minimum width [Hz] of filter.[default=100] ",
261
+ )
262
+ parser.add_argument(
263
+ "--maxBW",
264
+ type=int,
265
+ default=1000,
266
+ help="maximum width [Hz] of filter.[default=1000] ",
267
+ )
268
+ parser.add_argument(
269
+ "--minCoeff",
270
+ type=int,
271
+ default=10,
272
+ help="minimum filter coefficients. More the filter coefficients more ideal the filter slope.[default=10]",
273
+ )
274
+ parser.add_argument(
275
+ "--maxCoeff",
276
+ type=int,
277
+ default=100,
278
+ help="maximum filter coefficients. More the filter coefficients more ideal the filter slope.[default=100]",
279
+ )
280
+ parser.add_argument(
281
+ "--minG",
282
+ type=int,
283
+ default=0,
284
+ help="minimum gain factor of linear component.[default=0]",
285
+ )
286
+ parser.add_argument(
287
+ "--maxG",
288
+ type=int,
289
+ default=0,
290
+ help="maximum gain factor of linear component.[default=0]",
291
+ )
292
+ parser.add_argument(
293
+ "--minBiasLinNonLin",
294
+ type=int,
295
+ default=5,
296
+ help=" minimum gain difference between linear and non-linear components.[default=5]",
297
+ )
298
+ parser.add_argument(
299
+ "--maxBiasLinNonLin",
300
+ type=int,
301
+ default=20,
302
+ help=" maximum gain difference between linear and non-linear components.[default=20]",
303
+ )
304
+ parser.add_argument(
305
+ "--N_f",
306
+ type=int,
307
+ default=5,
308
+ help="order of the (non-)linearity where N_f=1 refers only to linear components.[default=5]",
309
+ )
310
+
311
+ # ISD_additive_noise parameters
312
+ parser.add_argument(
313
+ "--P",
314
+ type=int,
315
+ default=10,
316
+ help="Maximum number of uniformly distributed samples in [%].[defaul=10]",
317
+ )
318
+ parser.add_argument(
319
+ "--g_sd", type=int, default=2, help="gain parameters > 0. [default=2]"
320
+ )
321
+
322
+ # SSI_additive_noise parameters
323
+ parser.add_argument(
324
+ "--SNRmin",
325
+ type=int,
326
+ default=10,
327
+ help="Minimum SNR value for coloured additive noise.[defaul=10]",
328
+ )
329
+ parser.add_argument(
330
+ "--SNRmax",
331
+ type=int,
332
+ default=40,
333
+ help="Maximum SNR value for coloured additive noise.[defaul=40]",
334
+ )
335
+
336
+ ##===================================================Rawboost data augmentation ======================================================================#
337
+
338
+ load_dotenv()
339
+ wandb_api_key = os.getenv("WANDB_API_KEY")
340
+ wandb_project_name = os.getenv("WANDB_PROJECT_NAME")
341
+
342
+ if not os.path.exists("models"):
343
+ os.mkdir("models")
344
+ args = parser.parse_args()
345
+ wandb.login(key=wandb_api_key)
346
+ wandb.init(
347
+ project=wandb_project_name,
348
+ config={
349
+ "learning_rate": args.lr,
350
+ "epochs": args.num_epochs,
351
+ "batch_size": args.batch_size,
352
+ "weight_decay": args.weight_decay,
353
+ },
354
+ )
355
+
356
+ # make experiment reproducible
357
+ set_random_seed(args.seed, args)
358
+
359
+ # define model saving path
360
+ model_tag = "model_{}_{}_{}_{}".format(
361
+ args.loss, args.num_epochs, args.batch_size, args.lr
362
+ )
363
+ if args.comment:
364
+ model_tag = model_tag + "_{}".format(args.comment)
365
+ model_save_path = os.path.join(args.save_path, model_tag)
366
+
367
+ # set model save directory
368
+ if not os.path.exists(model_save_path):
369
+ os.mkdir(model_save_path)
370
+
371
+ # GPU device
372
+ device = "cuda" if torch.cuda.is_available() else "cpu"
373
+ print("Device: {}".format(device))
374
+
375
+ model = Model(args, device)
376
+ nb_params = sum([param.view(-1).size()[0] for param in model.parameters()])
377
+ model = model.to(device)
378
+ print("nb_params:", nb_params)
379
+
380
+ # set Adam optimizer
381
+ optimizer = torch.optim.Adam(
382
+ model.parameters(), lr=args.lr, weight_decay=args.weight_decay
383
+ )
384
+
385
+ if args.model_path:
386
+ model.load_state_dict(torch.load(args.model_path, map_location=device))
387
+ print("Model loaded : {}".format(args.model_path))
388
+
389
+ # evaluation
390
+
391
+ if args.eval:
392
+ file_eval = genSpoof_list(
393
+ dir_meta=args.test_list_path, is_train=False, is_eval=True
394
+ )
395
+ print("no. of eval trials", len(file_eval))
396
+ eval_set = Dataset_ASVspoof2021_eval(list_IDs=file_eval)
397
+ eval_output = os.path.join(
398
+ args.test_score_dir, f"{args.model_name}_model_score.txt"
399
+ )
400
+ produce_evaluation_file(
401
+ eval_set, model, device, eval_output, args.test_list_path
402
+ )
403
+ output_file = os.path.join(
404
+ args.test_score_dir, f"{args.model_name}_model_eer.txt"
405
+ )
406
+ eval_eer = calculate_tDCF_EER(
407
+ cm_scores_file=eval_output, output_file=output_file
408
+ )
409
+
410
+ sys.exit(0)
411
+
412
+ trn_list_path = args.trn_list_path
413
+ dev_trial_path = args.dev_list_path
414
+ train_set = Dataset_ASVspoof2019_train(args, metafile=trn_list_path, algo=args.algo)
415
+ train_loader = DataLoader(
416
+ train_set,
417
+ batch_size=args.batch_size,
418
+ num_workers=16,
419
+ shuffle=True,
420
+ drop_last=True,
421
+ )
422
+ del train_set
423
+
424
+ dev_set = Dataset_ASVspoof2019_train(args, metafile=dev_trial_path, algo=args.algo)
425
+ dev_loader = DataLoader(
426
+ dev_set, batch_size=args.batch_size, num_workers=16, shuffle=False
427
+ )
428
+ del dev_set
429
+ # Training and validation
430
+ num_epochs = args.num_epochs
431
+ writer = SummaryWriter("logs/{}".format(model_tag))
432
+
433
+ for epoch in range(num_epochs):
434
+
435
+ running_loss = train_epoch(
436
+ train_loader, model, args.lr, optimizer, device, args
437
+ )
438
+ val_loss = evaluate_accuracy(dev_loader, model, device, args)
439
+ wandb.log({"epoch": epoch, "train_loss": running_loss, "val_loss": val_loss})
440
+ writer.add_scalar("val_loss", val_loss, epoch)
441
+ writer.add_scalar("loss", running_loss, epoch)
442
+ print("\n{} - {} - {} ".format(epoch, running_loss, val_loss))
443
+ torch.save(
444
+ model.state_dict(),
445
+ os.path.join(model_save_path, "epoch_{}.pth".format(epoch)),
446
+ )
clean/image/aide/LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2024 Shilin Yan
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
clean/image/aide/README.md ADDED
@@ -0,0 +1,168 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ <div align="center">
2
+ <br>
3
+ <h3>A Sanity Check for AI-generated Image Detection</h3>
4
+
5
+ [Shilin Yan](https://scholar.google.com/citations?user=2VhjOykAAAAJ&hl=zh-CN&oi=ao)<sup>1†</sup>, Ouxiang Li<sup>1,2†</sup>, Jiayin Cai<sup>1†</sup>, [Yanbin Hao](https://scholar.google.com/citations?user=vhPSOkEAAAAJ&hl=en&oi=ao)<sup>2</sup>, [Xiaolong Jiang](https://scholar.google.com/citations?user=G0Ow8j8AAAAJ&hl=en&oi=ao)<sup>1</sup>, [Yao Hu](https://scholar.google.com/citations?user=LIu7k7wAAAAJ&hl=en)<sup>1</sup>, [Weidi Xie](https://scholar.google.com/citations?user=Vtrqj4gAAAAJ&hl=en)<sup>3‡</sup>
6
+
7
+ <div class="is-size-6 publication-authors">
8
+ <p class="footnote">
9
+ <span class="footnote-symbol"><sup>†</sup></span>Equal contribution
10
+ <span class="footnote-symbol"><sup>‡</sup></span>Corresponding author
11
+ </p>
12
+ </div>
13
+
14
+ <sup>1</sup>Xiaohongshu Inc. <sup>2</sup>University of Science and Technology of China <sup>3</sup>Shanghai Jiao Tong University
15
+
16
+
17
+ <p align="center">
18
+ <a href='https://shilinyan99.github.io/AIDE'>
19
+ <img src='https://img.shields.io/badge/Project-Page-pink?style=flat&logo=Google%20chrome&logoColor=pink'>
20
+ </a>
21
+ <a href='https://arxiv.org/abs/2406.19435'>
22
+ <img src='https://img.shields.io/badge/Arxiv-2406.19435-A42C25?style=flat&logo=arXiv&logoColor=A42C25'>
23
+ </a>
24
+ <a href='https://arxiv.org/pdf/2406.19435'>
25
+ <img src='https://img.shields.io/badge/Paper-PDF-yellow?style=flat&logo=arXiv&logoColor=yellow'>
26
+ </a>
27
+ <!-- <img src="https://visitor-badge.laobi.icu/badge?page_id=shilinyan99/AIDE" alt="visitors"> -->
28
+ </p>
29
+ </div>
30
+
31
+
32
+ <!-- <div align="center">
33
+ <h1>
34
+ <b>
35
+ A Sanity Check for AI-generated Image Detection
36
+ </b>
37
+ </h1>
38
+ </div> -->
39
+ ## 🔥 News
40
+ * [2025-01-23]🎉🎉🎉 AIDE is accepted by ICLR 2025.
41
+ * [2024-12-29]🔥🔥🔥 We release the Chamelon dataset.
42
+ * [2024-06-20]🔥🔥🔥 We release the code and checkpoints of AIDE.
43
+
44
+
45
+ ## 🔍 Chameleon
46
+
47
+
48
+ **License**:
49
+ ```
50
+ Chameleon is only used for academic research. Commercial use in any form is prohibited.
51
+ ```
52
+
53
+ 🌟🌟🌟 If you need the Chameleon dataset, please send an email to **tattoo.ysl@gmail.com**. 🔥🔥🔥
54
+
55
+
56
+
57
+ **Comparison of `Chameleon` with existing benchmarks.**
58
+
59
+ <p align="center"><img src="docs/Chameleon.jpg" width="800"/></p>
60
+
61
+ We visualize two contemporary AI-generated image benchmarks, namely:
62
+
63
+ - **(a) AIGCDetect Benchmark**
64
+ - **(b) GenImage Benchmark**
65
+
66
+ where all images are generated from publicly available generators, such as ProGAN (GAN-based), SD v1.4 (DM-based), and Midjourney (commercial API). These images are generated by unconditional situations or conditioned on simple prompts (e.g., *photo of a plane*) without delicate manual adjustments, thereby inclined to generate obvious artifacts in consistency and semantics (marked with <span style="color:red">red boxes</span>).
67
+
68
+ In contrast, our **`Chameleon`** dataset in **(c)** aims to simulate real-world scenarios by collecting diverse images from online websites, where these online images are carefully adjusted by photographers and AI artists.
69
+
70
+
71
+
72
+ ## 👀 Method
73
+
74
+ We conduct a sanity check on **"whether the task of AI-generated image detection has been solved"**. To start with, we present **Chameleon** dataset, consisting AI-generated images that are genuinely challenging for human perception. To quantify the generalization of existing methods, we evaluate 9 off-the-shelf AI-generated image detectors on **Chameleon** dataset. Upon analysis, almost all models classify AI-generated images as real ones. Later, we propose **AIDE**~(**A**I-generated **I**mage **DE**tector with Hybrid Features), which leverages multiple experts to simultaneously extract visual artifacts and noise patterns.
75
+
76
+ <p align="center"><img src="docs/network.png" width="800"/></p>
77
+
78
+
79
+
80
+
81
+ ## Requirements
82
+
83
+ We test the codes in the following environments, other versions may also be compatible:
84
+
85
+ - CUDA 11.8
86
+ - Python 3.10
87
+ - Pytorch 2.0.1
88
+
89
+
90
+ ## Setup
91
+
92
+ First, clone the repository locally.
93
+
94
+ ```
95
+ https://github.com/shilinyan99/AIDE
96
+ ```
97
+
98
+ Then, install Pytorch 2.0.1 using the conda environment.
99
+ ```
100
+ conda install pytorch==2.0.1 torchvision==0.15.2 torchaudio==2.0.2 -c pytorch
101
+ ```
102
+
103
+ Lastly, install the necessary packages and pycocotools.
104
+
105
+ ```
106
+ pip install -r requirements.txt
107
+ ```
108
+
109
+
110
+ ## Get Started
111
+
112
+ ### Training
113
+
114
+ ```
115
+ ./scripts/train.sh --data_path [/path/to/train_data] --eval_data_path [/path/to/eval_data] --resnet_path [/path/to/pretrained_resnet_path] --convnext_path [/path/to/pretrained_convnext_path] --output_dir [/path/to/output_dir] [other args]
116
+ ```
117
+
118
+ For example, training on ProGAN, run the following command:
119
+
120
+ ```
121
+ ./scripts/train.sh --data_path dataset/progan/train --eval_data_path dataset/progan/eval --resnet_path pretrained_ckpts/resnet50.pth --convnext_path pretrained_ckpts/open_clip_pytorch_model.bin --output_dir results/progan_train
122
+ ```
123
+
124
+ ### Inference
125
+
126
+ Inference using the trained model.
127
+ ```
128
+ ./scripts/eval.sh --data_path [/path/to/train_data] --eval_data_path [/path/to/eval_data] --resume [/path/to/progan_train] --eval True --output_dir [/path/to/output_dir]
129
+ ```
130
+
131
+ For example, evaluating the progan_train model, run the following command:
132
+
133
+ ```
134
+ ./scripts/eval.sh --data_path dataset/progan/train --eval_data_path dataset/progan/eval --resume results/progan_train/progan_train.pth --eval True --output_dir results/progan_train
135
+ ```
136
+
137
+
138
+
139
+ ## Dataset
140
+
141
+ ### Training Set
142
+ We adopt the training set in [CNNSpot](https://github.com/peterwang512/CNNDetection) and [GenImage](https://github.com/Andrew-Zhu/GenImage).
143
+
144
+ ### Test Set
145
+ The whole test set we used in our experiments can be downloaded from [AIGCDetectBenchmark](https://github.com/Ekko-zn/AIGCDetectBenchmark?tab=readme-ov-file) and [GenImage](https://github.com/Andrew-Zhu/GenImage).
146
+
147
+
148
+ ## Model Zoo
149
+
150
+ Our training checkpoints can be downloaded from [link](https://drive.google.com/drive/folders/1qx76UFvDpgCxaPLBCmsA2WY-SSzeJrd4?usp=sharing).
151
+
152
+ ## Acknowledgement
153
+
154
+ This repo is based on [ConvNeXt](https://github.com/facebookresearch/ConvNeXt-V2). We also refer to the repositories [CNNSpot](https://github.com/peterwang512/CNNDetection)、[AIGCDetectBenchmark](https://github.com/Ekko-zn/AIGCDetectBenchmark?tab=readme-ov-file)、[GenImage](https://github.com/Andrew-Zhu/GenImage) and [DNF](https://github.com/YichiCS/DNF). Thanks for their wonderful works.
155
+
156
+ ## Citation
157
+
158
+ ```
159
+ @article{yan2024sanity,
160
+ title={A Sanity Check for AI-generated Image Detection},
161
+ author={Yan, Shilin and Li, Ouxiang and Cai, Jiayin and Hao, Yanbin and Jiang, Xiaolong and Hu, Yao and Xie, Weidi},
162
+ journal={arXiv preprint arXiv:2406.19435},
163
+ year={2024}
164
+ }
165
+ ```
166
+
167
+ ## Contact
168
+ If you have any question about this project, please feel free to contact tattoo.ysl@gmail.com.
clean/image/aide/SOURCE.md ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Source: image/aide
2
+
3
+ | Field | Value |
4
+ |---|---|
5
+ | Upstream | https://github.com/shilinyan99/AIDE |
6
+ | Paper | https://arxiv.org/abs/2406.19435 |
7
+ | Commit SHA | **not recorded** -- the vendoring step did not preserve it |
8
+ | Mirrored on | 2026-09-15 |
9
+ | Upstream license | LICENSE |
10
+
11
+ This is a **mirror**, stripped to the files needed for inference. The full
12
+ untouched snapshot is at `archive/image__aide.tar.gz`.
13
+
14
+ This code is the work of its original authors and is **not** covered by the
15
+ DeepSafe project license. If you are an author and want this removed, open an
16
+ issue on https://github.com/deepsafehq/deepsafe-bench and it will be taken down
17
+ within 48 hours, no questions asked.
clean/image/aide/data/__init__.py ADDED
File without changes
clean/image/aide/data/dct.py ADDED
@@ -0,0 +1,107 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """DCT frequency decomposition module for AIDE.
2
+
3
+ Source: https://github.com/shilinyan99/AIDE (ICLR 2025)
4
+ Selects the most and least frequency-active image patches for
5
+ multi-view input to the AIDE deepfake detector.
6
+ """
7
+
8
+ import torch
9
+ import torch.nn as nn
10
+ import numpy as np
11
+
12
+
13
+ def DCT_mat(size):
14
+ m = [[(np.sqrt(1./size) if i == 0 else np.sqrt(2./size)) * np.cos((j + 0.5) * np.pi * i / size) for j in range(size)] for i in range(size)]
15
+ return m
16
+
17
+ def generate_filter(start, end, size):
18
+ return [[0. if i + j > end or i + j < start else 1. for j in range(size)] for i in range(size)]
19
+
20
+ def norm_sigma(x):
21
+ return 2. * torch.sigmoid(x) - 1.
22
+
23
+ class Filter(nn.Module):
24
+ def __init__(self, size, band_start, band_end, use_learnable=False, norm=False):
25
+ super(Filter, self).__init__()
26
+ self.use_learnable = use_learnable
27
+ self.base = nn.Parameter(torch.tensor(generate_filter(band_start, band_end, size)), requires_grad=False)
28
+ if self.use_learnable:
29
+ self.learnable = nn.Parameter(torch.randn(size, size), requires_grad=True)
30
+ self.learnable.data.normal_(0., 0.1)
31
+ self.norm = norm
32
+ if norm:
33
+ self.ft_num = nn.Parameter(torch.sum(torch.tensor(generate_filter(band_start, band_end, size))), requires_grad=False)
34
+
35
+ def forward(self, x):
36
+ if self.use_learnable:
37
+ filt = self.base + norm_sigma(self.learnable)
38
+ else:
39
+ filt = self.base
40
+ if self.norm:
41
+ y = x * filt / self.ft_num
42
+ else:
43
+ y = x * filt
44
+ return y
45
+
46
+ class DCT_base_Rec_Module(nn.Module):
47
+ def __init__(self, window_size=32, stride=16, output=256, grade_N=6, level_fliter=[0]):
48
+ super().__init__()
49
+ assert output % window_size == 0
50
+ assert len(level_fliter) > 0
51
+ self.window_size = window_size
52
+ self.grade_N = grade_N
53
+ self.level_N = len(level_fliter)
54
+ self.N = (output // window_size) * (output // window_size)
55
+ self._DCT_patch = nn.Parameter(torch.tensor(DCT_mat(window_size)).float(), requires_grad=False)
56
+ self._DCT_patch_T = nn.Parameter(torch.transpose(torch.tensor(DCT_mat(window_size)).float(), 0, 1), requires_grad=False)
57
+ self.unfold = nn.Unfold(kernel_size=(window_size, window_size), stride=stride)
58
+ self.fold0 = nn.Fold(output_size=(window_size, window_size), kernel_size=(window_size, window_size), stride=window_size)
59
+ level_f = [Filter(window_size, 0, window_size * 2)]
60
+ self.level_filters = nn.ModuleList([level_f[i] for i in level_fliter])
61
+ self.grade_filters = nn.ModuleList([Filter(window_size, window_size * 2. / grade_N * i, window_size * 2. / grade_N * (i+1), norm=True) for i in range(grade_N)])
62
+
63
+ def forward(self, x):
64
+ N = self.N
65
+ grade_N = self.grade_N
66
+ level_N = self.level_N
67
+ window_size = self.window_size
68
+ C, W, H = x.shape
69
+ x_unfold = self.unfold(x.unsqueeze(0)).squeeze(0)
70
+ _, L = x_unfold.shape
71
+ x_unfold = x_unfold.transpose(0, 1).reshape(L, C, window_size, window_size)
72
+ x_dct = self._DCT_patch @ x_unfold @ self._DCT_patch_T
73
+ y_list = []
74
+ for i in range(self.level_N):
75
+ x_pass = self.level_filters[i](x_dct)
76
+ y = self._DCT_patch_T @ x_pass @ self._DCT_patch
77
+ y_list.append(y)
78
+ level_x_unfold = torch.cat(y_list, dim=1)
79
+ grade = torch.zeros(L).to(x.device)
80
+ w, k = 1, 2
81
+ for _ in range(grade_N):
82
+ _x = torch.abs(x_dct)
83
+ _x = torch.log(_x + 1)
84
+ _x = self.grade_filters[_](_x)
85
+ _x = torch.sum(_x, dim=[1,2,3])
86
+ grade += w * _x
87
+ w *= k
88
+ _, idx = torch.sort(grade)
89
+ max_idx = torch.flip(idx, dims=[0])[:N]
90
+ maxmax_idx = max_idx[0]
91
+ maxmax_idx1 = max_idx[1] if len(max_idx) > 1 else max_idx[0]
92
+ min_idx = idx[:N]
93
+ minmin_idx = idx[0]
94
+ minmin_idx1 = idx[1] if len(min_idx) > 1 else idx[0]
95
+ x_minmin = torch.index_select(level_x_unfold, 0, minmin_idx)
96
+ x_maxmax = torch.index_select(level_x_unfold, 0, maxmax_idx)
97
+ x_minmin1 = torch.index_select(level_x_unfold, 0, minmin_idx1)
98
+ x_maxmax1 = torch.index_select(level_x_unfold, 0, maxmax_idx1)
99
+ x_minmin = x_minmin.reshape(1, level_N*C*window_size*window_size).transpose(0, 1)
100
+ x_maxmax = x_maxmax.reshape(1, level_N*C*window_size*window_size).transpose(0, 1)
101
+ x_minmin1 = x_minmin1.reshape(1, level_N*C*window_size*window_size).transpose(0, 1)
102
+ x_maxmax1 = x_maxmax1.reshape(1, level_N*C*window_size*window_size).transpose(0, 1)
103
+ x_minmin = self.fold0(x_minmin)
104
+ x_maxmax = self.fold0(x_maxmax)
105
+ x_minmin1 = self.fold0(x_minmin1)
106
+ x_maxmax1 = self.fold0(x_maxmax1)
107
+ return x_minmin, x_maxmax, x_minmin1, x_maxmax1
clean/image/aide/engine_finetune.py ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+
3
+ # All rights reserved.
4
+
5
+ # This source code is licensed under the license found in the
6
+ # LICENSE file in the root directory of this source tree.
7
+
8
+ import os
9
+ import math
10
+ from typing import Iterable, Optional
11
+
12
+ import torch
13
+ import torch.distributed as dist
14
+ from timm.data import Mixup
15
+ from timm.utils import accuracy, ModelEma
16
+
17
+ import utils
18
+ from utils import adjust_learning_rate
19
+ from scipy.special import softmax
20
+ from sklearn.metrics import (
21
+ average_precision_score,
22
+ accuracy_score
23
+ )
24
+ import numpy as np
25
+
26
+
27
+ def train_one_epoch(model: torch.nn.Module, criterion: torch.nn.Module,
28
+ data_loader: Iterable, optimizer: torch.optim.Optimizer,
29
+ device: torch.device, epoch: int, loss_scaler, max_norm: float = 0,
30
+ model_ema: Optional[ModelEma] = None, mixup_fn: Optional[Mixup] = None,
31
+ log_writer=None, args=None):
32
+ model.train(True)
33
+ metric_logger = utils.MetricLogger(delimiter=" ")
34
+ metric_logger.add_meter('lr', utils.SmoothedValue(window_size=1, fmt='{value:.6f}'))
35
+ header = 'Epoch: [{}]'.format(epoch)
36
+ print_freq = 20
37
+
38
+ update_freq = args.update_freq
39
+ use_amp = args.use_amp
40
+ optimizer.zero_grad()
41
+
42
+ for data_iter_step, (samples, targets) in enumerate(metric_logger.log_every(data_loader, print_freq, header)):
43
+ # we use a per iteration (instead of per epoch) lr scheduler
44
+ if data_iter_step % update_freq == 0:
45
+ adjust_learning_rate(optimizer, data_iter_step / len(data_loader) + epoch, args)
46
+
47
+ samples = samples.to(device, non_blocking=True)
48
+ targets = targets.to(device, non_blocking=True)
49
+
50
+ if mixup_fn is not None:
51
+ samples, targets = mixup_fn(samples, targets)
52
+
53
+ if use_amp:
54
+ with torch.cuda.amp.autocast():
55
+ output = model(samples)
56
+ loss = criterion(output, targets)
57
+ else: # full precision
58
+ output = model(samples)
59
+ loss = criterion(output, targets)
60
+
61
+ loss_value = loss.item()
62
+
63
+ if not math.isfinite(loss_value):
64
+ print("Loss is {}, stopping training".format(loss_value))
65
+ assert math.isfinite(loss_value)
66
+
67
+ if use_amp:
68
+ # this attribute is added by timm on one optimizer (adahessian)
69
+ is_second_order = hasattr(optimizer, 'is_second_order') and optimizer.is_second_order
70
+ loss /= update_freq
71
+ grad_norm = loss_scaler(loss, optimizer, clip_grad=max_norm,
72
+ parameters=model.parameters(), create_graph=is_second_order,
73
+ update_grad=(data_iter_step + 1) % update_freq == 0)
74
+ if (data_iter_step + 1) % update_freq == 0:
75
+ optimizer.zero_grad()
76
+ if model_ema is not None:
77
+ model_ema.update(model)
78
+ else: # full precision
79
+ loss /= update_freq
80
+ loss.backward()
81
+ if (data_iter_step + 1) % update_freq == 0:
82
+ optimizer.step()
83
+ optimizer.zero_grad()
84
+ if model_ema is not None:
85
+ model_ema.update(model)
86
+
87
+ torch.cuda.synchronize()
88
+
89
+ if mixup_fn is None:
90
+ class_acc = (output.max(-1)[-1] == targets).float().mean()
91
+ else:
92
+ class_acc = None
93
+
94
+ metric_logger.update(loss=loss_value)
95
+ metric_logger.update(class_acc=class_acc)
96
+ min_lr = 10.
97
+ max_lr = 0.
98
+ for group in optimizer.param_groups:
99
+ min_lr = min(min_lr, group["lr"])
100
+ max_lr = max(max_lr, group["lr"])
101
+
102
+ metric_logger.update(lr=max_lr)
103
+ metric_logger.update(min_lr=min_lr)
104
+ weight_decay_value = None
105
+ for group in optimizer.param_groups:
106
+ if group["weight_decay"] > 0:
107
+ weight_decay_value = group["weight_decay"]
108
+ metric_logger.update(weight_decay=weight_decay_value)
109
+ if use_amp:
110
+ metric_logger.update(grad_norm=grad_norm)
111
+ if log_writer is not None:
112
+ log_writer.update(loss=loss_value, head="loss")
113
+ log_writer.update(class_acc=class_acc, head="loss")
114
+ log_writer.update(lr=max_lr, head="opt")
115
+ log_writer.update(min_lr=min_lr, head="opt")
116
+ log_writer.update(weight_decay=weight_decay_value, head="opt")
117
+ if use_amp:
118
+ log_writer.update(grad_norm=grad_norm, head="opt")
119
+ log_writer.set_step()
120
+
121
+ # gather the stats from all processes
122
+ metric_logger.synchronize_between_processes()
123
+ print("Averaged stats:", metric_logger)
124
+ return {k: meter.global_avg for k, meter in metric_logger.meters.items()}
125
+
126
+ @torch.no_grad()
127
+ def evaluate(data_loader, model, device, use_amp=False):
128
+ criterion = torch.nn.CrossEntropyLoss()
129
+
130
+ metric_logger = utils.MetricLogger(delimiter=" ")
131
+ header = 'Test:'
132
+
133
+ # switch to evaluation mode
134
+ model.eval()
135
+
136
+ for index, batch in enumerate(metric_logger.log_every(data_loader, 10, header)):
137
+ images = batch[0]
138
+ target = batch[-1]
139
+
140
+ images = images.to(device, non_blocking=True)
141
+ target = target.to(device, non_blocking=True)
142
+
143
+ # compute output
144
+ if use_amp:
145
+ with torch.cuda.amp.autocast(dytpe=torch.bfloat16):
146
+ output = model(images)
147
+ if isinstance(output, dict):
148
+ output = output['logits']
149
+ loss = criterion(output, target)
150
+ else:
151
+ output = model(images) #[bs, num_cls]
152
+ if isinstance(output, dict):
153
+ output = output['logits']
154
+
155
+ loss = criterion(output, target)
156
+
157
+ if index == 0:
158
+ predictions = output
159
+ labels = target
160
+ else:
161
+ predictions = torch.cat((predictions, output), 0)
162
+ labels = torch.cat((labels, target), 0)
163
+
164
+ torch.cuda.synchronize()
165
+
166
+ acc1, acc5 = accuracy(output, target, topk=(1, 2))
167
+
168
+ batch_size = images.shape[0]
169
+ metric_logger.update(loss=loss.item())
170
+ metric_logger.meters['acc1'].update(acc1.item(), n=batch_size)
171
+ metric_logger.meters['acc5'].update(acc5.item(), n=batch_size)
172
+ # gather the stats from all processes
173
+ metric_logger.synchronize_between_processes()
174
+ print('* Acc@1 {top1.global_avg:.3f} Acc@5 {top5.global_avg:.3f} loss {losses.global_avg:.3f}'
175
+ .format(top1=metric_logger.acc1, top5=metric_logger.acc5, losses=metric_logger.loss))
176
+
177
+
178
+ output_ddp = [torch.zeros_like(predictions) for _ in range(utils.get_world_size())]
179
+ dist.all_gather(output_ddp, predictions)
180
+ labels_ddp = [torch.zeros_like(labels) for _ in range(utils.get_world_size())]
181
+ dist.all_gather(labels_ddp, labels)
182
+
183
+ output_all = torch.concat(output_ddp, dim=0)
184
+ labels_all = torch.concat(labels_ddp, dim=0)
185
+
186
+
187
+ y_pred = softmax(output_all.detach().cpu().numpy(), axis=1)[:, 1]
188
+ y_true = labels_all.detach().cpu().numpy()
189
+ y_true = y_true.astype(int)
190
+
191
+
192
+ acc = accuracy_score(y_true, y_pred > 0.5)
193
+ ap = average_precision_score(y_true, y_pred)
194
+
195
+
196
+
197
+ return {k: meter.global_avg for k, meter in metric_logger.meters.items()}, acc, ap
clean/image/aide/main_finetune.py ADDED
@@ -0,0 +1,449 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+
3
+ # All rights reserved.
4
+
5
+ # This source code is licensed under the license found in the
6
+ # LICENSE file in the root directory of this source tree.
7
+
8
+
9
+ import argparse
10
+ import datetime
11
+ import numpy as np
12
+ import time
13
+ import json
14
+ import os
15
+ from pathlib import Path
16
+
17
+ import torch
18
+ import torch.backends.cudnn as cudnn
19
+
20
+ from timm.models.layers import trunc_normal_
21
+ from timm.data.mixup import Mixup
22
+ from timm.loss import LabelSmoothingCrossEntropy, SoftTargetCrossEntropy
23
+ from timm.utils import ModelEma
24
+ from optim_factory import create_optimizer, LayerDecayValueAssigner
25
+
26
+ from data.datasets import TrainDataset, TestDataset
27
+ from engine_finetune import train_one_epoch, evaluate
28
+
29
+ import utils
30
+ from utils import NativeScalerWithGradNormCount as NativeScaler
31
+ from utils import str2bool, remap_checkpoint_keys
32
+ import models.AIDE as AIDE
33
+ import csv
34
+ import warnings
35
+
36
+ warnings.filterwarnings('ignore')
37
+
38
+ def get_args_parser():
39
+ parser = argparse.ArgumentParser('Resnet fine-tuning', add_help=False)
40
+ parser.add_argument('--batch_size', default=64, type=int,
41
+ help='Per GPU batch size')
42
+ parser.add_argument('--epochs', default=100, type=int)
43
+ parser.add_argument('--update_freq', default=1, type=int,
44
+ help='gradient accumulation steps')
45
+
46
+ # Model parameters
47
+ parser.add_argument('--model', default='AIDE', type=str, metavar='MODEL',
48
+ help='Name of model to train')
49
+ parser.add_argument('--resnet_path', default=None, type=str, metavar='MODEL',
50
+ help='Path of resnet model')
51
+ parser.add_argument('--convnext_path', default=None, type=str, metavar='MODEL',
52
+ help='Path of ConvNeXt of model ')
53
+
54
+ # EMA related parameters
55
+ parser.add_argument('--model_ema', type=str2bool, default=False)
56
+ parser.add_argument('--model_ema_decay', type=float, default=0.9999, help='')
57
+ parser.add_argument('--model_ema_force_cpu', type=str2bool, default=False, help='')
58
+ parser.add_argument('--model_ema_eval', type=str2bool, default=False, help='Using ema to eval during training.')
59
+
60
+ # Optimization parameters
61
+ parser.add_argument('--clip_grad', type=float, default=None, metavar='NORM',
62
+ help='Clip gradient norm (default: None, no clipping)')
63
+ parser.add_argument('--weight_decay', type=float, default=0.,
64
+ help='weight decay (default: 0.05)')
65
+ parser.add_argument('--lr', type=float, default=None, metavar='LR',
66
+ help='learning rate (absolute lr)')
67
+ parser.add_argument('--blr', type=float, default=5e-4, metavar='LR',
68
+ help='base learning rate: absolute_lr = base_lr * total_batch_size / 256')
69
+ parser.add_argument('--layer_decay', type=float, default=1.0)
70
+ parser.add_argument('--min_lr', type=float, default=1e-6, metavar='LR',
71
+ help='lower lr bound for cyclic schedulers that hit 0 (1e-6)')
72
+ parser.add_argument('--warmup_epochs', type=int, default=0, metavar='N',
73
+ help='epochs to warmup LR, if scheduler supports')
74
+
75
+ parser.add_argument('--warmup_steps', type=int, default=-1, metavar='N',
76
+ help='num of steps to warmup LR, will overload warmup_epochs if set > 0')
77
+ parser.add_argument('--opt', default='adamw', type=str, metavar='OPTIMIZER',
78
+ help='Optimizer (default: "adamw"')
79
+ parser.add_argument('--opt_eps', default=1e-8, type=float, metavar='EPSILON',
80
+ help='Optimizer Epsilon (default: 1e-8)')
81
+ parser.add_argument('--opt_betas', default=None, type=float, nargs='+', metavar='BETA',
82
+ help='Optimizer Betas (default: None, use opt default)')
83
+ parser.add_argument('--momentum', type=float, default=0.9, metavar='M',
84
+ help='SGD momentum (default: 0.9)')
85
+ parser.add_argument('--weight_decay_end', type=float, default=None, help="""Final value of the
86
+ weight decay. We use a cosine schedule for WD and using a larger decay by
87
+ the end of training improves performance for ViTs.""")
88
+
89
+ # Augmentation parameters
90
+ parser.add_argument('--color_jitter', type=float, default=None, metavar='PCT',
91
+ help='Color jitter factor (enabled only when not using Auto/RandAug)')
92
+ parser.add_argument('--aa', type=str, default='rand-m9-mstd0.5-inc1', metavar='NAME',
93
+ help='Use AutoAugment policy. "v0" or "original". " + "(default: rand-m9-mstd0.5-inc1)')
94
+ parser.add_argument('--smoothing', type=float, default=0.1,
95
+ help='Label smoothing (default: 0.1)')
96
+
97
+ parser.add_argument('--train_interpolation', type=str, default='bicubic',
98
+ help='Training interpolation (random, bilinear, bicubic default: "bicubic")')
99
+
100
+ # * Random Erase params
101
+ parser.add_argument('--reprob', type=float, default=0.25, metavar='PCT',
102
+ help='Random erase prob (default: 0.25)')
103
+ parser.add_argument('--remode', type=str, default='pixel',
104
+ help='Random erase mode (default: "pixel")')
105
+ parser.add_argument('--recount', type=int, default=1,
106
+ help='Random erase count (default: 1)')
107
+ parser.add_argument('--resplit', type=str2bool, default=False,
108
+ help='Do not random erase first (clean) augmentation split')
109
+
110
+ # * Mixup params
111
+ parser.add_argument('--mixup', type=float, default=0.,
112
+ help='mixup alpha, mixup enabled if > 0.')
113
+ parser.add_argument('--cutmix', type=float, default=0.,
114
+ help='cutmix alpha, cutmix enabled if > 0.')
115
+ parser.add_argument('--cutmix_minmax', type=float, nargs='+', default=None,
116
+ help='cutmix min/max ratio, overrides alpha and enables cutmix if set (default: None)')
117
+ parser.add_argument('--mixup_prob', type=float, default=1.0,
118
+ help='Probability of performing mixup or cutmix when either/both is enabled')
119
+ parser.add_argument('--mixup_switch_prob', type=float, default=0.5,
120
+ help='Probability of switching to cutmix when both mixup and cutmix enabled')
121
+ parser.add_argument('--mixup_mode', type=str, default='batch',
122
+ help='How to apply mixup/cutmix params. Per "batch", "pair", or "elem"')
123
+
124
+ # * Finetuning params
125
+ parser.add_argument('--finetune', default='',
126
+ help='finetune from checkpoint')
127
+ parser.add_argument('--head_init_scale', default=0.001, type=float,
128
+ help='classifier head initial scale, typically adjusted in fine-tuning')
129
+ parser.add_argument('--model_key', default='model|module', type=str,
130
+ help='which key to load from saved state dict, usually model or model_ema')
131
+ parser.add_argument('--model_prefix', default='', type=str)
132
+
133
+ # Dataset parameters
134
+ parser.add_argument('--data_path', default='path/dataset', type=str,
135
+ help='dataset path')
136
+ parser.add_argument('--nb_classes', default=2, type=int,
137
+ help='number of the classification types')
138
+ parser.add_argument('--output_dir', default='',
139
+ help='path where to save, empty for no saving')
140
+ parser.add_argument('--log_dir', default=None,
141
+ help='path where to tensorboard log')
142
+ parser.add_argument('--device', default='cuda',
143
+ help='device to use for training / testing')
144
+ parser.add_argument('--seed', default=0, type=int)
145
+ parser.add_argument('--resume', default='',
146
+ help='resume from checkpoint')
147
+
148
+ parser.add_argument('--eval_data_path', default=None, type=str,
149
+ help='dataset path for evaluation')
150
+ parser.add_argument('--imagenet_default_mean_and_std', type=str2bool, default=True)
151
+ parser.add_argument('--data_set', default='IMNET', choices=['CIFAR', 'IMNET', 'image_folder'],
152
+ type=str, help='ImageNet dataset path')
153
+ parser.add_argument('--auto_resume', type=str2bool, default=True)
154
+ parser.add_argument('--save_ckpt', type=str2bool, default=True)
155
+ parser.add_argument('--save_ckpt_freq', default=1, type=int)
156
+ parser.add_argument('--save_ckpt_num', default=100, type=int)
157
+
158
+ parser.add_argument('--start_epoch', default=0, type=int, metavar='N',
159
+ help='start epoch')
160
+ parser.add_argument('--eval', type=str2bool, default=False,
161
+ help='Perform evaluation only')
162
+ parser.add_argument('--dist_eval', type=str2bool, default=True,
163
+ help='Enabling distributed evaluation')
164
+ parser.add_argument('--disable_eval', type=str2bool, default=False,
165
+ help='Disabling evaluation during training')
166
+ parser.add_argument('--num_workers', default=16, type=int)
167
+ parser.add_argument('--pin_mem', type=str2bool, default=True,
168
+ help='Pin CPU memory in DataLoader for more efficient (sometimes) transfer to GPU.')
169
+
170
+ # Evaluation parameters
171
+ parser.add_argument('--crop_pct', type=float, default=None)
172
+
173
+ # distributed training parameters
174
+ parser.add_argument('--world_size', default=1, type=int,
175
+ help='number of distributed processes')
176
+ parser.add_argument('--local_rank', default=-1, type=int)
177
+ parser.add_argument('--dist_on_itp', type=str2bool, default=False)
178
+ parser.add_argument('--dist_url', default='env://',
179
+ help='url used to set up distributed training')
180
+
181
+ parser.add_argument('--use_amp', type=str2bool, default=False,
182
+ help="Use apex AMP (Automatic Mixed Precision) or not")
183
+ return parser
184
+
185
+ def main(args):
186
+ utils.init_distributed_mode(args)
187
+ print(args)
188
+ device = torch.device(args.device)
189
+
190
+ # fix the seed for reproducibility
191
+ seed = args.seed + utils.get_rank()
192
+ torch.manual_seed(seed)
193
+ np.random.seed(seed)
194
+ cudnn.benchmark = True
195
+
196
+ dataset_train = TrainDataset(is_train=True, args=args)
197
+
198
+ if args.disable_eval:
199
+ args.dist_eval = False
200
+ dataset_val = None
201
+ else:
202
+ dataset_val = TrainDataset(is_train=False, args=args)
203
+
204
+ num_tasks = utils.get_world_size()
205
+ global_rank = utils.get_rank()
206
+
207
+ sampler_train = torch.utils.data.DistributedSampler(
208
+ dataset_train, num_replicas=num_tasks, rank=global_rank, shuffle=True, seed=args.seed,
209
+ )
210
+ print("Sampler_train = %s" % str(sampler_train))
211
+ if args.dist_eval:
212
+ if len(dataset_val) % num_tasks != 0:
213
+ print('Warning: Enabling distributed evaluation with an eval dataset not divisible by process number. '
214
+ 'This will slightly alter validation results as extra duplicate entries are added to achieve '
215
+ 'equal num of samples per-process.')
216
+ sampler_val = torch.utils.data.DistributedSampler(
217
+ dataset_val, num_replicas=num_tasks, rank=global_rank, shuffle=False)
218
+ else:
219
+ sampler_val = torch.utils.data.SequentialSampler(dataset_val)
220
+
221
+ if global_rank == 0 and args.log_dir is not None:
222
+ os.makedirs(args.log_dir, exist_ok=True)
223
+ log_writer = utils.TensorboardLogger(log_dir=args.log_dir)
224
+ else:
225
+ log_writer = None
226
+
227
+ data_loader_train = torch.utils.data.DataLoader(
228
+ dataset_train, sampler=sampler_train,
229
+ batch_size=args.batch_size,
230
+ num_workers=args.num_workers,
231
+ pin_memory=args.pin_mem,
232
+ drop_last=True,
233
+ )
234
+ if dataset_val is not None:
235
+ data_loader_val = torch.utils.data.DataLoader(
236
+ dataset_val, sampler=sampler_val,
237
+ batch_size=args.batch_size,
238
+ num_workers=args.num_workers,
239
+ pin_memory=args.pin_mem,
240
+ drop_last=False
241
+ )
242
+ else:
243
+ data_loader_val = None
244
+
245
+ mixup_fn = None
246
+ mixup_active = args.mixup > 0 or args.cutmix > 0. or args.cutmix_minmax is not None
247
+ if mixup_active:
248
+ print("Mixup is activated!")
249
+ mixup_fn = Mixup(
250
+ mixup_alpha=args.mixup, cutmix_alpha=args.cutmix, cutmix_minmax=args.cutmix_minmax,
251
+ prob=args.mixup_prob, switch_prob=args.mixup_switch_prob, mode=args.mixup_mode,
252
+ label_smoothing=args.smoothing, num_classes=args.nb_classes)
253
+
254
+
255
+ model = AIDE.__dict__[args.model](
256
+ resnet_path=args.resnet_path,
257
+ convnext_path=args.convnext_path
258
+ )
259
+
260
+ model.to(device)
261
+
262
+ model_ema = None
263
+ if args.model_ema:
264
+ # Important to create EMA model after cuda(), DP wrapper, and AMP but before SyncBN and DDP wrapper
265
+ model_ema = ModelEma(
266
+ model,
267
+ decay=args.model_ema_decay,
268
+ device='cpu' if args.model_ema_force_cpu else '',
269
+ resume='')
270
+ print("Using EMA with decay = %.8f" % args.model_ema_decay)
271
+
272
+ model_without_ddp = model
273
+ n_parameters = sum(p.numel() for p in model.parameters() if p.requires_grad)
274
+
275
+ print("Model = %s" % str(model_without_ddp))
276
+ print('number of params:', n_parameters)
277
+
278
+ eff_batch_size = args.batch_size * args.update_freq * utils.get_world_size()
279
+ num_training_steps_per_epoch = len(dataset_train) // eff_batch_size
280
+
281
+ if args.lr is None:
282
+ args.lr = args.blr * eff_batch_size / 256
283
+
284
+ print("base lr: %.2e" % (args.lr * 256 / eff_batch_size))
285
+ print("actual lr: %.2e" % args.lr)
286
+
287
+ print("accumulate grad iterations: %d" % args.update_freq)
288
+ print("effective batch size: %d" % eff_batch_size)
289
+
290
+ assigner = None
291
+
292
+ if args.distributed:
293
+ model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[args.gpu], find_unused_parameters=True)
294
+ model_without_ddp = model.module
295
+
296
+ optimizer = create_optimizer(
297
+ args, model_without_ddp, skip_list=None,
298
+ get_num_layer=assigner.get_layer_id if assigner is not None else None,
299
+ get_layer_scale=assigner.get_scale if assigner is not None else None)
300
+ loss_scaler = NativeScaler()
301
+
302
+ if mixup_fn is not None:
303
+ # smoothing is handled with mixup label transform
304
+ criterion = SoftTargetCrossEntropy()
305
+ elif args.smoothing > 0.:
306
+ criterion = LabelSmoothingCrossEntropy(smoothing=args.smoothing)
307
+ else:
308
+ criterion = torch.nn.CrossEntropyLoss()
309
+
310
+ print("criterion = %s" % str(criterion))
311
+
312
+ utils.auto_load_model(
313
+ args=args, model=model, model_without_ddp=model_without_ddp,
314
+ optimizer=optimizer, loss_scaler=loss_scaler, model_ema=model_ema)
315
+
316
+ if args.eval:
317
+ print(f"Eval only mode")
318
+
319
+ vals = os.listdir(args.eval_data_path)
320
+ if len(vals) == 16:
321
+ vals = ["progan", "stylegan", "biggan", "cyclegan", "stargan", "gaugan", "stylegan2", "whichfaceisreal", "ADM", "Glide", "Midjourney", "stable_diffusion_v_1_4", "stable_diffusion_v_1_5", "VQDM", "wukong", "DALLE2"]
322
+ if len(vals) == 8:
323
+ vals = ["Midjourney", "stable_diffusion_v_1_4", "stable_diffusion_v_1_5", "ADM", "glide", "wukong", "VQDM", "BigGAN"]
324
+ eval_data_path = args.eval_data_path
325
+
326
+ rows = [["{} model testing on...".format(args.resume)],
327
+ ['testset', 'accuracy', 'avg precision']]
328
+
329
+ for v_id, val in enumerate(vals):
330
+
331
+ args.eval_data_path = os.path.join(args.eval_data_path, val)
332
+ dataset_val = TestDataset(is_train=False, args=args)
333
+ args.eval_data_path = eval_data_path
334
+
335
+ if args.dist_eval:
336
+ if len(dataset_val) % num_tasks != 0:
337
+ print('Warning: Enabling distributed evaluation with an eval dataset not divisible by process number. '
338
+ 'This will slightly alter validation results as extra duplicate entries are added to achieve '
339
+ 'equal num of samples per-process.')
340
+ sampler_val = torch.utils.data.DistributedSampler(
341
+ dataset_val, num_replicas=num_tasks, rank=global_rank, shuffle=False)
342
+ else:
343
+ sampler_val = torch.utils.data.SequentialSampler(dataset_val)
344
+
345
+ data_loader_val = torch.utils.data.DataLoader(
346
+ dataset_val, sampler=sampler_val,
347
+ batch_size=args.batch_size,
348
+ num_workers=args.num_workers,
349
+ pin_memory=args.pin_mem,
350
+ drop_last=False
351
+ )
352
+
353
+
354
+ test_stats, acc, ap = evaluate(data_loader_val, model, device)
355
+ print(f"Accuracy of the network on {len(dataset_val)} test images: {test_stats['acc1']:.5f}%")
356
+
357
+ print(f"test dataset is {val} acc: {acc}, ap: {ap}")
358
+ print("***********************************")
359
+
360
+ rows.append([val, acc, ap])
361
+
362
+
363
+ test_dataset_name = args.eval_data_path.split('/')[-2]
364
+
365
+ csv_name = os.path.join(args.output_dir, f'{os.path.basename(args.resume)}_{test_dataset_name}.csv')
366
+ with open(csv_name, 'w') as f:
367
+ csv_writer = csv.writer(f, delimiter=',')
368
+ csv_writer.writerows(rows)
369
+ return
370
+
371
+ max_accuracy = 0.0
372
+ if args.model_ema and args.model_ema_eval:
373
+ max_accuracy_ema = 0.0
374
+
375
+ print("Start training for %d epochs" % args.epochs)
376
+ start_time = time.time()
377
+ for epoch in range(args.start_epoch, args.epochs):
378
+ if args.distributed:
379
+ data_loader_train.sampler.set_epoch(epoch)
380
+ if log_writer is not None:
381
+ log_writer.set_step(epoch * num_training_steps_per_epoch * args.update_freq)
382
+ train_stats = train_one_epoch(
383
+ model, criterion, data_loader_train,
384
+ optimizer, device, epoch, loss_scaler,
385
+ args.clip_grad, model_ema, mixup_fn,
386
+ log_writer=log_writer,
387
+ args=args
388
+ )
389
+ if args.output_dir and args.save_ckpt:
390
+ if (epoch + 1) % args.save_ckpt_freq == 0 or epoch + 1 == args.epochs:
391
+ utils.save_model(
392
+ args=args, model=model, model_without_ddp=model_without_ddp, optimizer=optimizer,
393
+ loss_scaler=loss_scaler, epoch=epoch, model_ema=model_ema)
394
+ if data_loader_val is not None:
395
+ test_stats, acc, ap = evaluate(data_loader_val, model, device, use_amp=args.use_amp)
396
+ print(f"Accuracy of the model on the {len(dataset_val)} test images: {test_stats['acc1']:.1f}%, ap: {ap}.")
397
+ if max_accuracy < test_stats["acc1"]:
398
+ max_accuracy = test_stats["acc1"]
399
+ if args.output_dir and args.save_ckpt:
400
+ utils.save_model(
401
+ args=args, model=model, model_without_ddp=model_without_ddp, optimizer=optimizer,
402
+ loss_scaler=loss_scaler, epoch="best", model_ema=model_ema)
403
+ print(f'Max accuracy: {max_accuracy:.2f}%')
404
+
405
+ if log_writer is not None:
406
+ log_writer.update(test_acc1=test_stats['acc1'], head="perf", step=epoch)
407
+ log_writer.update(test_acc5=test_stats['acc5'], head="perf", step=epoch)
408
+ log_writer.update(test_loss=test_stats['loss'], head="perf", step=epoch)
409
+
410
+ log_stats = {**{f'train_{k}': v for k, v in train_stats.items()},
411
+ **{f'test_{k}': v for k, v in test_stats.items()},
412
+ 'epoch': epoch,
413
+ 'n_parameters': n_parameters}
414
+
415
+ # repeat testing routines for EMA, if ema eval is turned on
416
+ if args.model_ema and args.model_ema_eval:
417
+ test_stats_ema, acc, ap = evaluate(data_loader_val, model_ema.ema, device, use_amp=args.use_amp)
418
+ print(f"Accuracy of the model EMA on {len(dataset_val)} test images: {test_stats_ema['acc1']:.1f}%, ap: {ap}")
419
+ if max_accuracy_ema < test_stats_ema["acc1"]:
420
+ max_accuracy_ema = test_stats_ema["acc1"]
421
+ if args.output_dir and args.save_ckpt:
422
+ utils.save_model(
423
+ args=args, model=model, model_without_ddp=model_without_ddp, optimizer=optimizer,
424
+ loss_scaler=loss_scaler, epoch="best-ema", model_ema=model_ema)
425
+ print(f'Max EMA accuracy: {max_accuracy_ema:.2f}%')
426
+ if log_writer is not None:
427
+ log_writer.update(test_acc1_ema=test_stats_ema['acc1'], head="perf", step=epoch)
428
+ log_stats.update({**{f'test_{k}_ema': v for k, v in test_stats_ema.items()}})
429
+ else:
430
+ log_stats = {**{f'train_{k}': v for k, v in train_stats.items()},
431
+ 'epoch': epoch,
432
+ 'n_parameters': n_parameters}
433
+
434
+ if args.output_dir and utils.is_main_process():
435
+ if log_writer is not None:
436
+ log_writer.flush()
437
+ with open(os.path.join(args.output_dir, "log.txt"), mode="a", encoding="utf-8") as f:
438
+ f.write(json.dumps(log_stats) + "\n")
439
+
440
+ total_time = time.time() - start_time
441
+ total_time_str = str(datetime.timedelta(seconds=int(total_time)))
442
+ print('Training time {}'.format(total_time_str))
443
+
444
+ if __name__ == '__main__':
445
+ parser = argparse.ArgumentParser('AIDE traning', parents=[get_args_parser()])
446
+ args = parser.parse_args()
447
+ if args.output_dir:
448
+ Path(args.output_dir).mkdir(parents=True, exist_ok=True)
449
+ main(args)
clean/image/aide/models/AIDE.py ADDED
@@ -0,0 +1,298 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch.nn as nn
2
+ import torch.utils.model_zoo as model_zoo
3
+ import torch
4
+ import clip
5
+ import open_clip
6
+ from .srm_filter_kernel import all_normalized_hpf_list
7
+ import numpy as np
8
+
9
+ class HPF(nn.Module):
10
+ def __init__(self):
11
+ super(HPF, self).__init__()
12
+
13
+ #Load 30 SRM Filters
14
+ all_hpf_list_5x5 = []
15
+
16
+ for hpf_item in all_normalized_hpf_list:
17
+ if hpf_item.shape[0] == 3:
18
+ hpf_item = np.pad(hpf_item, pad_width=((1, 1), (1, 1)), mode='constant')
19
+
20
+ all_hpf_list_5x5.append(hpf_item)
21
+
22
+ hpf_weight = torch.Tensor(all_hpf_list_5x5).view(30, 1, 5, 5).contiguous()
23
+ hpf_weight = torch.nn.Parameter(hpf_weight.repeat(1, 3, 1, 1), requires_grad=False)
24
+
25
+
26
+ self.hpf = nn.Conv2d(3, 30, kernel_size=5, padding=2, bias=False)
27
+ self.hpf.weight = hpf_weight
28
+
29
+
30
+ def forward(self, input):
31
+
32
+ output = self.hpf(input)
33
+
34
+ return output
35
+
36
+
37
+
38
+ def conv3x3(in_planes, out_planes, stride=1):
39
+ """3x3 convolution with padding"""
40
+ return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,
41
+ padding=1, bias=False)
42
+
43
+
44
+ def conv1x1(in_planes, out_planes, stride=1):
45
+ """1x1 convolution"""
46
+ return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)
47
+
48
+
49
+ class BasicBlock(nn.Module):
50
+ expansion = 1
51
+
52
+ def __init__(self, inplanes, planes, stride=1, downsample=None):
53
+ super(BasicBlock, self).__init__()
54
+ self.conv1 = conv3x3(inplanes, planes, stride)
55
+ self.bn1 = nn.BatchNorm2d(planes)
56
+ self.relu = nn.ReLU(inplace=True)
57
+ self.conv2 = conv3x3(planes, planes)
58
+ self.bn2 = nn.BatchNorm2d(planes)
59
+ self.downsample = downsample
60
+ self.stride = stride
61
+
62
+ def forward(self, x):
63
+ identity = x
64
+
65
+ out = self.conv1(x)
66
+ out = self.bn1(out)
67
+ out = self.relu(out)
68
+
69
+ out = self.conv2(out)
70
+ out = self.bn2(out)
71
+
72
+ if self.downsample is not None:
73
+ identity = self.downsample(x)
74
+
75
+ out += identity
76
+ out = self.relu(out)
77
+
78
+ return out
79
+
80
+
81
+ class Bottleneck(nn.Module):
82
+ expansion = 4
83
+
84
+ def __init__(self, inplanes, planes, stride=1, downsample=None):
85
+ super(Bottleneck, self).__init__()
86
+ self.conv1 = conv1x1(inplanes, planes)
87
+ self.bn1 = nn.BatchNorm2d(planes)
88
+ self.conv2 = conv3x3(planes, planes, stride)
89
+ self.bn2 = nn.BatchNorm2d(planes)
90
+ self.conv3 = conv1x1(planes, planes * self.expansion)
91
+ self.bn3 = nn.BatchNorm2d(planes * self.expansion)
92
+ self.relu = nn.ReLU(inplace=True)
93
+ self.downsample = downsample
94
+ self.stride = stride
95
+
96
+ def forward(self, x):
97
+ identity = x
98
+
99
+ out = self.conv1(x)
100
+ out = self.bn1(out)
101
+ out = self.relu(out)
102
+
103
+ out = self.conv2(out)
104
+ out = self.bn2(out)
105
+ out = self.relu(out)
106
+
107
+ out = self.conv3(out)
108
+ out = self.bn3(out)
109
+
110
+ if self.downsample is not None:
111
+ identity = self.downsample(x)
112
+
113
+ out += identity
114
+ out = self.relu(out)
115
+
116
+ return out
117
+
118
+
119
+ class ResNet(nn.Module):
120
+
121
+ def __init__(self, block, layers, num_classes=1000, zero_init_residual=True):
122
+ super(ResNet, self).__init__()
123
+
124
+ self.inplanes = 64
125
+ self.conv1 = nn.Conv2d(30, 64, kernel_size=7, stride=2, padding=3,
126
+ bias=False)
127
+ self.bn1 = nn.BatchNorm2d(64)
128
+ self.relu = nn.ReLU(inplace=True)
129
+ self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
130
+ self.layer1 = self._make_layer(block, 64, layers[0])
131
+ self.layer2 = self._make_layer(block, 128, layers[1], stride=2)
132
+ self.layer3 = self._make_layer(block, 256, layers[2], stride=2)
133
+ self.layer4 = self._make_layer(block, 512, layers[3], stride=2)
134
+ self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
135
+ self.fc = nn.Linear(512 * block.expansion, num_classes)
136
+
137
+ for m in self.modules():
138
+ if isinstance(m, nn.Conv2d):
139
+ nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
140
+ elif isinstance(m, nn.BatchNorm2d):
141
+ nn.init.constant_(m.weight, 1)
142
+ nn.init.constant_(m.bias, 0)
143
+
144
+ # Zero-initialize the last BN in each residual branch,
145
+ # so that the residual branch starts with zeros, and each residual block behaves like an identity.
146
+ # This improves the model by 0.2~0.3% according to https://arxiv.org/abs/1706.02677
147
+ if zero_init_residual:
148
+ for m in self.modules():
149
+ if isinstance(m, Bottleneck):
150
+ nn.init.constant_(m.bn3.weight, 0)
151
+ elif isinstance(m, BasicBlock):
152
+ nn.init.constant_(m.bn2.weight, 0)
153
+
154
+ def _make_layer(self, block, planes, blocks, stride=1):
155
+ downsample = None
156
+ if stride != 1 or self.inplanes != planes * block.expansion:
157
+ downsample = nn.Sequential(
158
+ conv1x1(self.inplanes, planes * block.expansion, stride),
159
+ nn.BatchNorm2d(planes * block.expansion),
160
+ )
161
+
162
+ layers = []
163
+ layers.append(block(self.inplanes, planes, stride, downsample))
164
+ self.inplanes = planes * block.expansion
165
+ for _ in range(1, blocks):
166
+ layers.append(block(self.inplanes, planes))
167
+
168
+ return nn.Sequential(*layers)
169
+
170
+ def forward(self, x):
171
+
172
+ x = self.conv1(x)
173
+ x = self.bn1(x)
174
+ x = self.relu(x)
175
+ x = self.maxpool(x)
176
+
177
+ x = self.layer1(x)
178
+ x = self.layer2(x)
179
+ x = self.layer3(x)
180
+ x = self.layer4(x)
181
+
182
+ x = self.avgpool(x)
183
+ x = x.view(x.size(0), -1)
184
+
185
+
186
+ return x
187
+
188
+ class Mlp(nn.Module):
189
+ """ MLP as used in Vision Transformer, MLP-Mixer and related networks
190
+ """
191
+
192
+ def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU):
193
+ super().__init__()
194
+ out_features = out_features or in_features
195
+ hidden_features = hidden_features or in_features
196
+
197
+ self.fc1 = nn.Linear(in_features, hidden_features)
198
+ self.act = act_layer()
199
+ self.fc2 = nn.Linear(hidden_features, out_features)
200
+
201
+ def forward(self, x):
202
+ x = self.fc1(x)
203
+ x = self.act(x)
204
+ x = self.fc2(x)
205
+ return x
206
+
207
+ class AIDE_Model(nn.Module):
208
+
209
+ def __init__(self, resnet_path, convnext_path):
210
+ super(AIDE_Model, self).__init__()
211
+ self.hpf = HPF()
212
+ self.model_min = ResNet(Bottleneck, [3, 4, 6, 3])
213
+ self.model_max = ResNet(Bottleneck, [3, 4, 6, 3])
214
+
215
+ if resnet_path is not None:
216
+ pretrained_dict = torch.load(resnet_path, map_location='cpu')
217
+
218
+ model_min_dict = self.model_min.state_dict()
219
+ model_max_dict = self.model_max.state_dict()
220
+
221
+ for k in pretrained_dict.keys():
222
+ if k in model_min_dict and pretrained_dict[k].size() == model_min_dict[k].size():
223
+ model_min_dict[k] = pretrained_dict[k]
224
+ model_max_dict[k] = pretrained_dict[k]
225
+ else:
226
+ print(f"Skipping layer {k} because of size mismatch")
227
+
228
+ self.fc = Mlp(2048 + 256 , 1024, 2)
229
+
230
+ print("build model with convnext_xxl")
231
+ self.openclip_convnext_xxl, _, _ = open_clip.create_model_and_transforms(
232
+ "convnext_xxlarge", pretrained=convnext_path
233
+ )
234
+
235
+ self.openclip_convnext_xxl = self.openclip_convnext_xxl.visual.trunk
236
+ self.openclip_convnext_xxl.head.global_pool = nn.Identity()
237
+ self.openclip_convnext_xxl.head.flatten = nn.Identity()
238
+
239
+ self.openclip_convnext_xxl.eval()
240
+
241
+ self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
242
+ self.convnext_proj = nn.Sequential(
243
+ nn.Linear(3072, 256),
244
+
245
+ )
246
+ for param in self.openclip_convnext_xxl.parameters():
247
+ param.requires_grad = False
248
+
249
+
250
+
251
+ def forward(self, x):
252
+
253
+ b, t, c, h, w = x.shape
254
+
255
+ x_minmin = x[:, 0] #[b, c, h, w]
256
+ x_maxmax = x[:, 1]
257
+ x_minmin1 = x[:, 2]
258
+ x_maxmax1 = x[:, 3]
259
+ tokens = x[:, 4]
260
+
261
+ x_minmin = self.hpf(x_minmin)
262
+ x_maxmax = self.hpf(x_maxmax)
263
+ x_minmin1 = self.hpf(x_minmin1)
264
+ x_maxmax1 = self.hpf(x_maxmax1)
265
+
266
+ with torch.no_grad():
267
+
268
+ clip_mean = torch.Tensor([0.48145466, 0.4578275, 0.40821073])
269
+ clip_mean = clip_mean.to(tokens, non_blocking=True).view(3, 1, 1)
270
+ clip_std = torch.Tensor([0.26862954, 0.26130258, 0.27577711])
271
+ clip_std = clip_std.to(tokens, non_blocking=True).view(3, 1, 1)
272
+ dinov2_mean = torch.Tensor([0.485, 0.456, 0.406]).to(tokens, non_blocking=True).view(3, 1, 1)
273
+ dinov2_std = torch.Tensor([0.229, 0.224, 0.225]).to(tokens, non_blocking=True).view(3, 1, 1)
274
+
275
+ local_convnext_image_feats = self.openclip_convnext_xxl(
276
+ tokens * (dinov2_std / clip_std) + (dinov2_mean - clip_mean) / clip_std
277
+ ) #[b, 3072, 8, 8]
278
+ assert local_convnext_image_feats.size()[1:] == (3072, 8, 8)
279
+ local_convnext_image_feats = self.avgpool(local_convnext_image_feats).view(tokens.size(0), -1)
280
+ x_0 = self.convnext_proj(local_convnext_image_feats)
281
+
282
+ x_min = self.model_min(x_minmin)
283
+ x_max = self.model_max(x_maxmax)
284
+ x_min1 = self.model_min(x_minmin1)
285
+ x_max1 = self.model_max(x_maxmax1)
286
+
287
+ x_1 = (x_min + x_max + x_min1 + x_max1) / 4
288
+
289
+ x = torch.cat([x_0, x_1], dim=1)
290
+
291
+ x = self.fc(x)
292
+
293
+ return x
294
+
295
+ def AIDE(resnet_path, convnext_path):
296
+ model = AIDE_Model(resnet_path, convnext_path)
297
+ return model
298
+
clean/image/aide/models/__init__.py ADDED
File without changes
clean/image/aide/models/srm_filter_kernel.py ADDED
@@ -0,0 +1,220 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import numpy as np
3
+
4
+ filter_class_1 = [
5
+ np.array([
6
+ [1, 0, 0],
7
+ [0, -1, 0],
8
+ [0, 0, 0]
9
+ ], dtype=np.float32),
10
+ np.array([
11
+ [0, 1, 0],
12
+ [0, -1, 0],
13
+ [0, 0, 0]
14
+ ], dtype=np.float32),
15
+ np.array([
16
+ [0, 0, 1],
17
+ [0, -1, 0],
18
+ [0, 0, 0]
19
+ ], dtype=np.float32),
20
+ np.array([
21
+ [0, 0, 0],
22
+ [1, -1, 0],
23
+ [0, 0, 0]
24
+ ], dtype=np.float32),
25
+ np.array([
26
+ [0, 0, 0],
27
+ [0, -1, 1],
28
+ [0, 0, 0]
29
+ ], dtype=np.float32),
30
+ np.array([
31
+ [0, 0, 0],
32
+ [0, -1, 0],
33
+ [1, 0, 0]
34
+ ], dtype=np.float32),
35
+ np.array([
36
+ [0, 0, 0],
37
+ [0, -1, 0],
38
+ [0, 1, 0]
39
+ ], dtype=np.float32),
40
+ np.array([
41
+ [0, 0, 0],
42
+ [0, -1, 0],
43
+ [0, 0, 1]
44
+ ], dtype=np.float32)
45
+ ]
46
+
47
+
48
+ filter_class_2 = [
49
+ np.array([
50
+ [1, 0, 0],
51
+ [0, -2, 0],
52
+ [0, 0, 1]
53
+ ], dtype=np.float32),
54
+ np.array([
55
+ [0, 1, 0],
56
+ [0, -2, 0],
57
+ [0, 1, 0]
58
+ ], dtype=np.float32),
59
+ np.array([
60
+ [0, 0, 1],
61
+ [0, -2, 0],
62
+ [1, 0, 0]
63
+ ], dtype=np.float32),
64
+ np.array([
65
+ [0, 0, 0],
66
+ [1, -2, 1],
67
+ [0, 0, 0]
68
+ ], dtype=np.float32),
69
+ ]
70
+
71
+
72
+ filter_class_3 = [
73
+ np.array([
74
+ [-1, 0, 0, 0, 0],
75
+ [0, 3, 0, 0, 0],
76
+ [0, 0, -3, 0, 0],
77
+ [0, 0, 0, 1, 0],
78
+ [0, 0, 0, 0, 0]
79
+ ], dtype=np.float32),
80
+ np.array([
81
+ [0, 0, -1, 0, 0],
82
+ [0, 0, 3, 0, 0],
83
+ [0, 0, -3, 0, 0],
84
+ [0, 0, 1, 0, 0],
85
+ [0, 0, 0, 0, 0]
86
+ ], dtype=np.float32),
87
+ np.array([
88
+ [0, 0, 0, 0, -1],
89
+ [0, 0, 0, 3, 0],
90
+ [0, 0, -3, 0, 0],
91
+ [0, 1, 0, 0, 0],
92
+ [0, 0, 0, 0, 0]
93
+ ], dtype=np.float32),
94
+ np.array([
95
+ [0, 0, 0, 0, 0],
96
+ [0, 0, 0, 0, 0],
97
+ [0, 1, -3, 3, -1],
98
+ [0, 0, 0, 0, 0],
99
+ [0, 0, 0, 0, 0]
100
+ ], dtype=np.float32),
101
+ np.array([
102
+ [0, 0, 0, 0, 0],
103
+ [0, 1, 0, 0, 0],
104
+ [0, 0, -3, 0, 0],
105
+ [0, 0, 0, 3, 0],
106
+ [0, 0, 0, 0, -1]
107
+ ], dtype=np.float32),
108
+ np.array([
109
+ [0, 0, 0, 0, 0],
110
+ [0, 0, 1, 0, 0],
111
+ [0, 0, -3, 0, 0],
112
+ [0, 0, 3, 0, 0],
113
+ [0, 0, -1, 0, 0]
114
+ ], dtype=np.float32),
115
+ np.array([
116
+ [0, 0, 0, 0, 0],
117
+ [0, 0, 0, 1, 0],
118
+ [0, 0, -3, 0, 0],
119
+ [0, 3, 0, 0, 0],
120
+ [-1, 0, 0, 0, 0]
121
+ ], dtype=np.float32),
122
+ np.array([
123
+ [0, 0, 0, 0, 0],
124
+ [0, 0, 0, 0, 0],
125
+ [-1, 3, -3, 1, 0],
126
+ [0, 0, 0, 0, 0],
127
+ [0, 0, 0, 0, 0]
128
+ ], dtype=np.float32)
129
+ ]
130
+
131
+
132
+ filter_edge_3x3 = [
133
+ np.array([
134
+ [-1, 2, -1],
135
+ [2, -4, 2],
136
+ [0, 0, 0]
137
+ ], dtype=np.float32),
138
+ np.array([
139
+ [0, 2, -1],
140
+ [0, -4, 2],
141
+ [0, 2, -1]
142
+ ], dtype=np.float32),
143
+ np.array([
144
+ [0, 0, 0],
145
+ [2, -4, 2],
146
+ [-1, 2, -1]
147
+ ], dtype=np.float32),
148
+ np.array([
149
+ [-1, 2, 0],
150
+ [2, -4, 0],
151
+ [-1, 2, 0]
152
+ ], dtype=np.float32),
153
+ ]
154
+
155
+ filter_edge_5x5 = [
156
+ np.array([
157
+ [-1, 2, -2, 2, -1],
158
+ [2, -6, 8, -6, 2],
159
+ [-2, 8, -12, 8, -2],
160
+ [0, 0, 0, 0, 0],
161
+ [0, 0, 0, 0, 0]
162
+ ], dtype=np.float32),
163
+ np.array([
164
+ [0, 0, -2, 2, -1],
165
+ [0, 0, 8, -6, 2],
166
+ [0, 0, -12, 8, -2],
167
+ [0, 0, 8, -6, 2],
168
+ [0, 0, -2, 2, -1]
169
+ ], dtype=np.float32),
170
+ np.array([
171
+ [0, 0, 0, 0, 0],
172
+ [0, 0, 0, 0, 0],
173
+ [-2, 8, -12, 8, -2],
174
+ [2, -6, 8, -6, 2],
175
+ [-1, 2, -2, 2, -1]
176
+ ], dtype=np.float32),
177
+ np.array([
178
+ [-1, 2, -2, 0, 0],
179
+ [2, -6, 8, 0, 0],
180
+ [-2, 8, -12, 0, 0],
181
+ [2, -6, 8, 0, 0],
182
+ [-1, 2, -2, 0, 0]
183
+ ], dtype=np.float32),
184
+ ]
185
+
186
+ square_3x3 = np.array([
187
+ [-1, 2, -1],
188
+ [2, -4, 2],
189
+ [-1, 2, -1]
190
+ ], dtype=np.float32)
191
+
192
+ square_5x5 = np.array([
193
+ [-1, 2, -2, 2, -1],
194
+ [2, -6, 8, -6, 2],
195
+ [-2, 8, -12, 8, -2],
196
+ [2, -6, 8, -6, 2],
197
+ [-1, 2, -2, 2, -1]
198
+ ], dtype=np.float32)
199
+
200
+
201
+ all_hpf_list = filter_class_1 + filter_class_2 + filter_class_3 + filter_edge_3x3 + filter_edge_5x5 + [square_3x3, square_5x5]
202
+
203
+ hpf_3x3_list = filter_class_1 + filter_class_2 + filter_edge_3x3 + [square_3x3]
204
+ hpf_5x5_list = filter_class_3 + filter_edge_5x5 + [square_5x5]
205
+
206
+ normalized_filter_class_2 = [hpf / 2 for hpf in filter_class_2]
207
+ normalized_filter_class_3 = [hpf / 3 for hpf in filter_class_3]
208
+ normalized_filter_edge_3x3 = [hpf / 4 for hpf in filter_edge_3x3]
209
+ normalized_square_3x3 = square_3x3 / 4
210
+ normalized_filter_edge_5x5 = [hpf / 12 for hpf in filter_edge_5x5]
211
+ normalized_square_5x5 = square_5x5 / 12
212
+
213
+ all_normalized_hpf_list = filter_class_1 + normalized_filter_class_2 + normalized_filter_class_3 + \
214
+ normalized_filter_edge_3x3 + normalized_filter_edge_5x5 + [normalized_square_3x3, normalized_square_5x5]
215
+
216
+ normalized_hpf_3x3_list = filter_class_1 + normalized_filter_class_2 + normalized_filter_edge_3x3 + [normalized_square_3x3]
217
+ normalized_hpf_5x5_list = normalized_filter_class_3 + normalized_filter_edge_5x5 + [normalized_square_5x5]
218
+
219
+ normalized_3x3_list = normalized_filter_edge_3x3 + [normalized_square_3x3]
220
+ normalized_5x5_list = normalized_filter_edge_5x5 + [normalized_square_5x5]
clean/image/aide/models/utils.py ADDED
@@ -0,0 +1,116 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+
3
+ # All rights reserved.
4
+
5
+ # This source code is licensed under the license found in the
6
+ # LICENSE file in the root directory of this source tree.
7
+
8
+
9
+ import numpy.random as random
10
+
11
+ import torch
12
+ import torch.nn as nn
13
+ import torch.nn.functional as F
14
+ # from MinkowskiEngine import SparseTensor
15
+
16
+ # class MinkowskiGRN(nn.Module):
17
+ # """ GRN layer for sparse tensors.
18
+ # """
19
+ # def __init__(self, dim):
20
+ # super().__init__()
21
+ # self.gamma = nn.Parameter(torch.zeros(1, dim))
22
+ # self.beta = nn.Parameter(torch.zeros(1, dim))
23
+
24
+ # def forward(self, x):
25
+ # cm = x.coordinate_manager
26
+ # in_key = x.coordinate_map_key
27
+
28
+ # Gx = torch.norm(x.F, p=2, dim=0, keepdim=True)
29
+ # Nx = Gx / (Gx.mean(dim=-1, keepdim=True) + 1e-6)
30
+ # return SparseTensor(
31
+ # self.gamma * (x.F * Nx) + self.beta + x.F,
32
+ # coordinate_map_key=in_key,
33
+ # coordinate_manager=cm)
34
+
35
+ # class MinkowskiDropPath(nn.Module):
36
+ # """ Drop Path for sparse tensors.
37
+ # """
38
+
39
+ # def __init__(self, drop_prob: float = 0., scale_by_keep: bool = True):
40
+ # super(MinkowskiDropPath, self).__init__()
41
+ # self.drop_prob = drop_prob
42
+ # self.scale_by_keep = scale_by_keep
43
+
44
+ # def forward(self, x):
45
+ # if self.drop_prob == 0. or not self.training:
46
+ # return x
47
+ # cm = x.coordinate_manager
48
+ # in_key = x.coordinate_map_key
49
+ # keep_prob = 1 - self.drop_prob
50
+ # mask = torch.cat([
51
+ # torch.ones(len(_)) if random.uniform(0, 1) > self.drop_prob
52
+ # else torch.zeros(len(_)) for _ in x.decomposed_coordinates
53
+ # ]).view(-1, 1).to(x.device)
54
+ # if keep_prob > 0.0 and self.scale_by_keep:
55
+ # mask.div_(keep_prob)
56
+ # return SparseTensor(
57
+ # x.F * mask,
58
+ # coordinate_map_key=in_key,
59
+ # coordinate_manager=cm)
60
+
61
+ # class MinkowskiLayerNorm(nn.Module):
62
+ # """ Channel-wise layer normalization for sparse tensors.
63
+ # """
64
+
65
+ # def __init__(
66
+ # self,
67
+ # normalized_shape,
68
+ # eps=1e-6,
69
+ # ):
70
+ # super(MinkowskiLayerNorm, self).__init__()
71
+ # self.ln = nn.LayerNorm(normalized_shape, eps=eps)
72
+ # def forward(self, input):
73
+ # output = self.ln(input.F)
74
+ # return SparseTensor(
75
+ # output,
76
+ # coordinate_map_key=input.coordinate_map_key,
77
+ # coordinate_manager=input.coordinate_manager)
78
+
79
+ class LayerNorm(nn.Module):
80
+ """ LayerNorm that supports two data formats: channels_last (default) or channels_first.
81
+ The ordering of the dimensions in the inputs. channels_last corresponds to inputs with
82
+ shape (batch_size, height, width, channels) while channels_first corresponds to inputs
83
+ with shape (batch_size, channels, height, width).
84
+ """
85
+ def __init__(self, normalized_shape, eps=1e-6, data_format="channels_last"):
86
+ super().__init__()
87
+ self.weight = nn.Parameter(torch.ones(normalized_shape))
88
+ self.bias = nn.Parameter(torch.zeros(normalized_shape))
89
+ self.eps = eps
90
+ self.data_format = data_format
91
+ if self.data_format not in ["channels_last", "channels_first"]:
92
+ raise NotImplementedError
93
+ self.normalized_shape = (normalized_shape, )
94
+
95
+ def forward(self, x):
96
+ if self.data_format == "channels_last":
97
+ return F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps)
98
+ elif self.data_format == "channels_first":
99
+ u = x.mean(1, keepdim=True)
100
+ s = (x - u).pow(2).mean(1, keepdim=True)
101
+ x = (x - u) / torch.sqrt(s + self.eps)
102
+ x = self.weight[:, None, None] * x + self.bias[:, None, None]
103
+ return x
104
+
105
+ class GRN(nn.Module):
106
+ """ GRN (Global Response Normalization) layer
107
+ """
108
+ def __init__(self, dim):
109
+ super().__init__()
110
+ self.gamma = nn.Parameter(torch.zeros(1, 1, 1, dim))
111
+ self.beta = nn.Parameter(torch.zeros(1, 1, 1, dim))
112
+
113
+ def forward(self, x):
114
+ Gx = torch.norm(x, p=2, dim=(1,2), keepdim=True)
115
+ Nx = Gx / (Gx.mean(dim=-1, keepdim=True) + 1e-6)
116
+ return self.gamma * (x * Nx) + self.beta + x
clean/image/aide/optim_factory.py ADDED
@@ -0,0 +1,222 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates.
2
+
3
+ # All rights reserved.
4
+
5
+ # This source code is licensed under the license found in the
6
+ # LICENSE file in the root directory of this source tree.
7
+
8
+
9
+ import torch
10
+ from torch import optim as optim
11
+
12
+ from timm.optim.adafactor import Adafactor
13
+ from timm.optim.adahessian import Adahessian
14
+ from timm.optim.adamp import AdamP
15
+ from timm.optim.lookahead import Lookahead
16
+ from timm.optim.nadam import Nadam
17
+ # from timm.optim.novograd import NovoGrad
18
+ # from timm.optim.nvnovograd import NvNovoGrad
19
+ from timm.optim.radam import RAdam
20
+ from timm.optim.rmsprop_tf import RMSpropTF
21
+ from timm.optim.sgdp import SGDP
22
+
23
+ import json
24
+
25
+ try:
26
+ from apex.optimizers import FusedNovoGrad, FusedAdam, FusedLAMB, FusedSGD
27
+ has_apex = True
28
+ except ImportError:
29
+ has_apex = False
30
+
31
+
32
+ def get_num_layer_for_convnext_single(var_name, depths):
33
+ """
34
+ Each layer is assigned distinctive layer ids
35
+ """
36
+ if var_name.startswith("downsample_layers"):
37
+ stage_id = int(var_name.split('.')[1])
38
+ layer_id = sum(depths[:stage_id]) + 1
39
+ return layer_id
40
+
41
+ elif var_name.startswith("stages"):
42
+ stage_id = int(var_name.split('.')[1])
43
+ block_id = int(var_name.split('.')[2])
44
+ layer_id = sum(depths[:stage_id]) + block_id + 1
45
+ return layer_id
46
+
47
+ else:
48
+ return sum(depths) + 1
49
+
50
+
51
+ def get_num_layer_for_convnext(var_name):
52
+ """
53
+ Divide [3, 3, 27, 3] layers into 12 groups; each group is three
54
+ consecutive blocks, including possible neighboring downsample layers;
55
+ adapted from https://github.com/microsoft/unilm/blob/master/beit/optim_factory.py
56
+ """
57
+ num_max_layer = 12
58
+ if var_name.startswith("downsample_layers"):
59
+ stage_id = int(var_name.split('.')[1])
60
+ if stage_id == 0:
61
+ layer_id = 0
62
+ elif stage_id == 1 or stage_id == 2:
63
+ layer_id = stage_id + 1
64
+ elif stage_id == 3:
65
+ layer_id = 12
66
+ return layer_id
67
+
68
+ elif var_name.startswith("stages"):
69
+ stage_id = int(var_name.split('.')[1])
70
+ block_id = int(var_name.split('.')[2])
71
+ if stage_id == 0 or stage_id == 1:
72
+ layer_id = stage_id + 1
73
+ elif stage_id == 2:
74
+ layer_id = 3 + block_id // 3
75
+ elif stage_id == 3:
76
+ layer_id = 12
77
+ return layer_id
78
+ else:
79
+ return num_max_layer + 1
80
+
81
+ class LayerDecayValueAssigner(object):
82
+ def __init__(self, values, depths=[3,3,27,3], layer_decay_type='single'):
83
+ self.values = values
84
+ self.depths = depths
85
+ self.layer_decay_type = layer_decay_type
86
+
87
+ def get_scale(self, layer_id):
88
+ return self.values[layer_id]
89
+
90
+ def get_layer_id(self, var_name):
91
+ if self.layer_decay_type == 'single':
92
+ return get_num_layer_for_convnext_single(var_name, self.depths)
93
+ else:
94
+ return get_num_layer_for_convnext(var_name)
95
+
96
+
97
+ def get_parameter_groups(model, weight_decay=1e-5, skip_list=(), get_num_layer=None, get_layer_scale=None):
98
+ parameter_group_names = {}
99
+ parameter_group_vars = {}
100
+
101
+ for name, param in model.named_parameters():
102
+ if not param.requires_grad:
103
+ continue # frozen weights
104
+ if len(param.shape) == 1 or name.endswith(".bias") or name in skip_list or \
105
+ name.endswith(".gamma") or name.endswith(".beta"):
106
+ group_name = "no_decay"
107
+ this_weight_decay = 0.
108
+ else:
109
+ group_name = "decay"
110
+ this_weight_decay = weight_decay
111
+ if get_num_layer is not None:
112
+ layer_id = get_num_layer(name)
113
+ group_name = "layer_%d_%s" % (layer_id, group_name)
114
+ else:
115
+ layer_id = None
116
+
117
+ if group_name not in parameter_group_names:
118
+ if get_layer_scale is not None:
119
+ scale = get_layer_scale(layer_id)
120
+ else:
121
+ scale = 1.
122
+
123
+ parameter_group_names[group_name] = {
124
+ "weight_decay": this_weight_decay,
125
+ "params": [],
126
+ "lr_scale": scale
127
+ }
128
+ parameter_group_vars[group_name] = {
129
+ "weight_decay": this_weight_decay,
130
+ "params": [],
131
+ "lr_scale": scale
132
+ }
133
+
134
+ parameter_group_vars[group_name]["params"].append(param)
135
+ parameter_group_names[group_name]["params"].append(name)
136
+ print("Param groups = %s" % json.dumps(parameter_group_names, indent=2))
137
+ return list(parameter_group_vars.values())
138
+
139
+
140
+ def create_optimizer(args, model, get_num_layer=None, get_layer_scale=None, filter_bias_and_bn=True, skip_list=None):
141
+ opt_lower = args.opt.lower()
142
+ weight_decay = args.weight_decay
143
+ # if weight_decay and filter_bias_and_bn:
144
+ if filter_bias_and_bn:
145
+ skip = {}
146
+ if skip_list is not None:
147
+ skip = skip_list
148
+ elif hasattr(model, 'no_weight_decay'):
149
+ skip = model.no_weight_decay()
150
+ parameters = get_parameter_groups(model, weight_decay, skip, get_num_layer, get_layer_scale)
151
+ weight_decay = 0.
152
+ else:
153
+ parameters = model.parameters()
154
+
155
+ if 'fused' in opt_lower:
156
+ assert has_apex and torch.cuda.is_available(), 'APEX and CUDA required for fused optimizers'
157
+
158
+ opt_args = dict(lr=args.lr, weight_decay=weight_decay)
159
+ if hasattr(args, 'opt_eps') and args.opt_eps is not None:
160
+ opt_args['eps'] = args.opt_eps
161
+ if hasattr(args, 'opt_betas') and args.opt_betas is not None:
162
+ opt_args['betas'] = args.opt_betas
163
+
164
+ opt_split = opt_lower.split('_')
165
+ opt_lower = opt_split[-1]
166
+ if opt_lower == 'sgd' or opt_lower == 'nesterov':
167
+ opt_args.pop('eps', None)
168
+ optimizer = optim.SGD(parameters, momentum=args.momentum, nesterov=True, **opt_args)
169
+ elif opt_lower == 'momentum':
170
+ opt_args.pop('eps', None)
171
+ optimizer = optim.SGD(parameters, momentum=args.momentum, nesterov=False, **opt_args)
172
+ elif opt_lower == 'adam':
173
+ optimizer = optim.Adam(parameters, **opt_args)
174
+ elif opt_lower == 'adamw':
175
+ optimizer = optim.AdamW(parameters, **opt_args)
176
+ elif opt_lower == 'nadam':
177
+ optimizer = Nadam(parameters, **opt_args)
178
+ elif opt_lower == 'radam':
179
+ optimizer = RAdam(parameters, **opt_args)
180
+ elif opt_lower == 'adamp':
181
+ optimizer = AdamP(parameters, wd_ratio=0.01, nesterov=True, **opt_args)
182
+ elif opt_lower == 'sgdp':
183
+ optimizer = SGDP(parameters, momentum=args.momentum, nesterov=True, **opt_args)
184
+ elif opt_lower == 'adadelta':
185
+ optimizer = optim.Adadelta(parameters, **opt_args)
186
+ elif opt_lower == 'adafactor':
187
+ if not args.lr:
188
+ opt_args['lr'] = None
189
+ optimizer = Adafactor(parameters, **opt_args)
190
+ elif opt_lower == 'adahessian':
191
+ optimizer = Adahessian(parameters, **opt_args)
192
+ elif opt_lower == 'rmsprop':
193
+ optimizer = optim.RMSprop(parameters, alpha=0.9, momentum=args.momentum, **opt_args)
194
+ elif opt_lower == 'rmsproptf':
195
+ optimizer = RMSpropTF(parameters, alpha=0.9, momentum=args.momentum, **opt_args)
196
+ elif opt_lower == 'novograd':
197
+ optimizer = NovoGrad(parameters, **opt_args)
198
+ elif opt_lower == 'nvnovograd':
199
+ optimizer = NvNovoGrad(parameters, **opt_args)
200
+ elif opt_lower == 'fusedsgd':
201
+ opt_args.pop('eps', None)
202
+ optimizer = FusedSGD(parameters, momentum=args.momentum, nesterov=True, **opt_args)
203
+ elif opt_lower == 'fusedmomentum':
204
+ opt_args.pop('eps', None)
205
+ optimizer = FusedSGD(parameters, momentum=args.momentum, nesterov=False, **opt_args)
206
+ elif opt_lower == 'fusedadam':
207
+ optimizer = FusedAdam(parameters, adam_w_mode=False, **opt_args)
208
+ elif opt_lower == 'fusedadamw':
209
+ optimizer = FusedAdam(parameters, adam_w_mode=True, **opt_args)
210
+ elif opt_lower == 'fusedlamb':
211
+ optimizer = FusedLAMB(parameters, **opt_args)
212
+ elif opt_lower == 'fusednovograd':
213
+ opt_args.setdefault('betas', (0.95, 0.98))
214
+ optimizer = FusedNovoGrad(parameters, **opt_args)
215
+ else:
216
+ assert False and "Invalid optimizer"
217
+
218
+ if len(opt_split) > 1:
219
+ if opt_split[0] == 'lookahead':
220
+ optimizer = Lookahead(optimizer)
221
+
222
+ return optimizer
clean/image/aide/requirements.txt ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ einops==0.6.1
2
+ fairscale==0.4.13
3
+ filelock==3.13.1
4
+ ftfy==6.1.3
5
+ h5py==3.10.0
6
+ imgaug==0.2.6
7
+ keras==2.11.0
8
+ kornia==0.7.2
9
+ kornia_rs==0.1.2
10
+ lmdb==1.4.1
11
+ matplotlib==3.7.4
12
+ matplotlib-inline==0.1.6
13
+ numpy==1.24.3
14
+ omegaconf==2.3.0
15
+ open-clip-torch==2.24.0
16
+ openai-clip==1.0.1
17
+ openpyxl==3.1.2
18
+ pandas==2.0.3
19
+ Pillow==9.5.0
20
+ safetensors==0.4.1
21
+ scikit-image==0.20.0
22
+ scikit-learn==1.3.2
23
+ scipy==1.9.1
24
+ sentencepiece==0.2.0
25
+ streamlit==1.30.0
26
+ tenacity==8.2.3
27
+ tensorboard==2.11.2
28
+ tensorboard-data-server==0.6.1
29
+ tensorboard-plugin-wit==1.8.1
30
+ tensorboardX==2.6.2.2
31
+ timm==0.9.6
32
+ torch==1.11.0
33
+ torch-fidelity==0.3.0
34
+ torchmetrics==0.6.0
35
+ torchsummary==1.5.1
36
+ torchvision==0.12.0
37
+ tqdm==4.66.1