Add stripped inference-only model code mirror
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- clean/audio/nes2net/SOURCE.md +17 -0
- clean/audio/nes2net/__init__.py +0 -0
- clean/audio/nes2net/wav2vec2_Nes2Net_X.py +317 -0
- clean/audio/safeear/.gitignore +167 -0
- clean/audio/safeear/LICENSE +23 -0
- clean/audio/safeear/README.md +134 -0
- clean/audio/safeear/SOURCE.md +17 -0
- clean/audio/safeear/config/train19.yaml +87 -0
- clean/audio/safeear/config/train21.yaml +87 -0
- clean/audio/safeear/requirements.txt +123 -0
- clean/audio/safeear/safeear/losses/loss.py +215 -0
- clean/audio/safeear/safeear/models/decouple.py +207 -0
- clean/audio/safeear/safeear/models/discriminator.py +422 -0
- clean/audio/safeear/safeear/models/modules/__init__.py +21 -0
- clean/audio/safeear/safeear/models/modules/conv.py +252 -0
- clean/audio/safeear/safeear/models/modules/lstm.py +32 -0
- clean/audio/safeear/safeear/models/modules/norm.py +28 -0
- clean/audio/safeear/safeear/models/modules/quantization/__init__.py +8 -0
- clean/audio/safeear/safeear/models/modules/quantization/ac.py +292 -0
- clean/audio/safeear/safeear/models/modules/quantization/core_vq.py +366 -0
- clean/audio/safeear/safeear/models/modules/quantization/distrib.py +126 -0
- clean/audio/safeear/safeear/models/modules/quantization/vq.py +108 -0
- clean/audio/safeear/safeear/models/modules/seanet.py +275 -0
- clean/audio/safeear/safeear/models/safeear.py +959 -0
- clean/audio/safeear/safeear/trainer/safeear_trainer.py +189 -0
- clean/audio/safeear/safeear/utils/dump_hubert_feature.py +108 -0
- clean/audio/safeear/test.py +78 -0
- clean/audio/safeear/train.py +102 -0
- clean/audio/shiftyspeech/.env +2 -0
- clean/audio/shiftyspeech/LICENSE +21 -0
- clean/audio/shiftyspeech/RawBoost.py +143 -0
- clean/audio/shiftyspeech/SOURCE.md +17 -0
- clean/audio/shiftyspeech/Simplified_CM_solution.py +227 -0
- clean/audio/shiftyspeech/data_utils.py +292 -0
- clean/audio/shiftyspeech/model.py +603 -0
- clean/audio/shiftyspeech/startup_config.py +60 -0
- clean/audio/shiftyspeech/train.py +446 -0
- clean/image/aide/LICENSE +21 -0
- clean/image/aide/README.md +168 -0
- clean/image/aide/SOURCE.md +17 -0
- clean/image/aide/data/__init__.py +0 -0
- clean/image/aide/data/dct.py +107 -0
- clean/image/aide/engine_finetune.py +197 -0
- clean/image/aide/main_finetune.py +449 -0
- clean/image/aide/models/AIDE.py +298 -0
- clean/image/aide/models/__init__.py +0 -0
- clean/image/aide/models/srm_filter_kernel.py +220 -0
- clean/image/aide/models/utils.py +116 -0
- clean/image/aide/optim_factory.py +222 -0
- 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 |
+
[](https://arxiv.org/abs/2409.09272)
|
| 4 |
+
[](https://makeapullrequest.com)
|
| 5 |
+
[](https://creativecommons.org/licenses/by/4.0/)
|
| 6 |
+

|
| 7 |
+

|
| 8 |
+

|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
By [1] Zhejiang University, [2] Tsinghua University.
|
| 12 |
+
* [Xinfeng Li](https://letterligo.github.io)* [1], [Kai Li](https://cslikai.cn)* [2], Yifan Zheng [1], Chen Yan† [1], Xiaoyu Ji [1], Wenyuan Xu [1].
|
| 13 |
+
|
| 14 |
+
This repository is an official implementation of the SafeEar accepted to **ACM CCS 2024** (Core-A*, CCF-A, Big4) .
|
| 15 |
+
|
| 16 |
+
Please also visit our <a href="https://safeearweb.github.io/Project/">(1) Project Website</a>, <a href="https://zenodo.org/records/14062964">(2) Full CVoiceFake Dataset</a>, and <a href="https://zenodo.org/records/11124319">(3) Sampled CVoiceFake Dataset</a>.
|
| 17 |
+
|
| 18 |
+
## 🔥News
|
| 19 |
+
|
| 20 |
+
[2025-03-18]: Supported the batch testing for ASVspoof 2019 and 2021, fixed some bugs for datasets and trainer.
|
| 21 |
+
|
| 22 |
+
[2024-12-10]: Fixed all the bugs for training and test, and uploaded the files for data generation `datas/`.
|
| 23 |
+
|
| 24 |
+
[2024-12-01]: Uploaded the checkpoint for data generation `datas/`.
|
| 25 |
+
|
| 26 |
+
## ✨Key Highlights:
|
| 27 |
+
|
| 28 |
+
In this paper, we propose SafeEar, a novel framework that aims to detect deepfake audios without relying on accessing the speech content within. Our key idea is to devise a neural audio codec into a novel decoupling model that well separates the semantic and acoustic information from audio samples, and only use the acoustic information (e.g., prosody and timbre) for deepfake detection. In this way, no semantic content will be exposed to the detector. To overcome the challenge of identifying diverse deepfake audio without semantic clues, we enhance our deepfake detector with multi-head self-attention and codec augmentation. Extensive experiments conducted on four benchmark datasets demonstrate SafeEar’s effectiveness in detecting various deepfake techniques with an equal error rate (EER) down to 2.02%. Simultaneously, it shields five-language speech content from being deciphered by both machine and human auditory analysis, demonstrated by word error rates (WERs) all above 93.93% and our user study. Furthermore, our benchmark constructed for anti-deepfake and anti-content recovery evaluation helps provide a basis for future research in the realms of audio privacy preservation and deepfake detection.
|
| 29 |
+
|
| 30 |
+
## 🚀Overall Pipeline
|
| 31 |
+
|
| 32 |
+

|
| 33 |
+
|
| 34 |
+
## 🔧Installation
|
| 35 |
+
|
| 36 |
+
1. Clone the repository:
|
| 37 |
+
|
| 38 |
+
```shell
|
| 39 |
+
git clone git@github.com:LetterLiGo/SafeEar.git
|
| 40 |
+
cd SafeEar/
|
| 41 |
+
```
|
| 42 |
+
|
| 43 |
+
2. Create and activate the conda environment:
|
| 44 |
+
|
| 45 |
+
```shell
|
| 46 |
+
conda create -n safeear python=3.9
|
| 47 |
+
conda activate safeear
|
| 48 |
+
```
|
| 49 |
+
|
| 50 |
+
3. Install PyTorch and torchvision following the [official instructions](https://pytorch.org). The code requires `python=3.9`, `pytorch=1.13`, `torchvision=0.14`.
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
```shell
|
| 54 |
+
pip install torch==1.13.1+cu116 torchvision==0.14.1+cu116 torchaudio==0.13.1 --extra-index-url https://download.pytorch.org/whl/cu116
|
| 55 |
+
|
| 56 |
+
```
|
| 57 |
+
4. Install other dependencies:
|
| 58 |
+
|
| 59 |
+
```shell
|
| 60 |
+
pip install pip==24.0
|
| 61 |
+
pip install -r requirements.txt
|
| 62 |
+
```
|
| 63 |
+
|
| 64 |
+
## 📊Model Performance
|
| 65 |
+
### ASVspoof 2019 & 2021
|
| 66 |
+

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

|
| 69 |
+
|
| 70 |
+
## Data preparation
|
| 71 |
+
|
| 72 |
+
### AVSpoof 2019 & 2021
|
| 73 |
+
|
| 74 |
+
Please download the [ASVspoof 2019](https://datashare.is.ed.ac.uk/handle/10283/3336) and [ASVspoof 2021](https://www.asvspoof.org/index2021.html) datasets and extract them to the `datas/datasets` directory.
|
| 75 |
+
|
| 76 |
+
```shell
|
| 77 |
+
datas/datasets/ASVspoof2019
|
| 78 |
+
datas/datasets/ASVspoof2021
|
| 79 |
+
```
|
| 80 |
+
|
| 81 |
+
#### Generate the Hubert L9 feature files
|
| 82 |
+
|
| 83 |
+
```shell
|
| 84 |
+
mkdir model_zoos
|
| 85 |
+
cd model_zoos
|
| 86 |
+
wget https://dl.fbaipublicfiles.com/hubert/hubert_base_ls960.pt
|
| 87 |
+
wget https://cloud.tsinghua.edu.cn/f/413a0cd2e6f749eea956/?dl=1 -O SpeechTokenizer.pt
|
| 88 |
+
cd ../datas
|
| 89 |
+
# Generate the Hubert L9 feature files for ASVspoof 2019
|
| 90 |
+
python dump_hubert_avg_feature.py datasets/ASVSpoof2019 datasets/ASVSpoof2019_Hubert_L9
|
| 91 |
+
# Generate the Hubert L9 feature files for ASVspoof 2021
|
| 92 |
+
python dump_hubert_avg_feature.py datasets/ASVSpoof2021 datasets/ASVSpoof2021_Hubert_L9
|
| 93 |
+
```
|
| 94 |
+
|
| 95 |
+
## 📚Training
|
| 96 |
+
|
| 97 |
+
Before starting training, please modify the parameter configurations in [`configs`](configs).
|
| 98 |
+
|
| 99 |
+
Use the following commands to start training:
|
| 100 |
+
|
| 101 |
+
```shell
|
| 102 |
+
python train.py --conf_dir config/train19.yaml
|
| 103 |
+
python train.py --conf_dir config/train21.yaml
|
| 104 |
+
```
|
| 105 |
+
|
| 106 |
+
## 📈Testing/Inference
|
| 107 |
+
|
| 108 |
+
To evaluate a model on one or more GPUs, specify the `CUDA_VISIBLE_DEVICES`, `dataset`, `model` and `checkpoint`:
|
| 109 |
+
|
| 110 |
+
```shell
|
| 111 |
+
python test.py --conf_dir Exps/ASVspoof19/config.yaml
|
| 112 |
+
python test.py --conf_dir Exps/ASVspoof21/config.yaml
|
| 113 |
+
```
|
| 114 |
+
|
| 115 |
+
## Bugs and Issues
|
| 116 |
+
|
| 117 |
+
If you meet `RuntimeError: Failed to load audio from <_io.BytesIO object at 0x7f45cb978f90>`, please use the following command to fix it:
|
| 118 |
+
|
| 119 |
+
```shell
|
| 120 |
+
conda install -c anaconda 'ffmpeg<4.4'
|
| 121 |
+
```
|
| 122 |
+
|
| 123 |
+
## 📜Citation
|
| 124 |
+
|
| 125 |
+
If you find our work/code/dataset helpful, please consider citing:
|
| 126 |
+
|
| 127 |
+
```
|
| 128 |
+
@inproceedings{li2024safeear,
|
| 129 |
+
author = {Li, Xinfeng and Li, Kai and Zheng, Yifan and Yan, Chen and Ji, Xiaoyu, and Xu, Wenyuan},
|
| 130 |
+
title = {{SafeEar: Content Privacy-Preserving Audio Deepfake Detection}},
|
| 131 |
+
booktitle = {Proceedings of the 2024 {ACM} {SIGSAC} Conference on Computer and Communications Security (CCS)}
|
| 132 |
+
year = {2024},
|
| 133 |
+
}
|
| 134 |
+
```
|
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
|