Spaces:
No application file
No application file
File size: 3,623 Bytes
a5deb1f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 | import torch
import torch.nn as nn
from models.cpc import load_CPC, get_cnn_layer
"""
Encoder should take a wave file, then return an embedding
"""
class CPC_encoder(nn.Module):
def __init__(self):
super().__init__()
self.sample_rate = 16000
self.encoder = load_CPC(True)
self.output_dim = self.encoder.gEncoder.conv4.out_channels
self.dim = self.output_dim
self.downsample_ratio = 256
self.downsample = get_cnn_layer(
dim=self.output_dim,
kernel=[4],
stride=[4],
dilation=[1],
activation="GELU",
)
def forward(self, waveform):
if waveform.ndim < 3:
waveform = waveform.unsqueeze(1)
z = self.encoder.gEncoder(waveform)
z = z.permute(0, 2, 1)
#z = self.encoder.gAR(z)
z = self.downsample(z)
return z
class Audio_Block(nn.Module):
def __init__(self, in_channels, out_channels):
super(Audio_Block, self).__init__()
self.relu = nn.ReLU()
self.m_3 = nn.Conv2d(in_channels, out_channels, kernel_size = (3, 1), padding = (1, 0), bias = False)
self.bn_m_3 = nn.BatchNorm2d(out_channels, momentum = 0.01, eps = 0.001)
self.t_3 = nn.Conv2d(out_channels, out_channels, kernel_size = (1, 3), padding = (0, 1), bias = False)
self.bn_t_3 = nn.BatchNorm2d(out_channels, momentum = 0.01, eps = 0.001)
self.m_5 = nn.Conv2d(in_channels, out_channels, kernel_size = (5, 1), padding = (2, 0), bias = False)
self.bn_m_5 = nn.BatchNorm2d(out_channels, momentum = 0.01, eps = 0.001)
self.t_5 = nn.Conv2d(out_channels, out_channels, kernel_size = (1, 5), padding = (0, 2), bias = False)
self.bn_t_5 = nn.BatchNorm2d(out_channels, momentum = 0.01, eps = 0.001)
self.last = nn.Conv2d(out_channels, out_channels, kernel_size = (1, 1), padding = (0, 0), bias = False)
self.bn_last = nn.BatchNorm2d(out_channels, momentum = 0.01, eps = 0.001)
def forward(self, x):
x_3 = self.relu(self.bn_m_3(self.m_3(x)))
x_3 = self.relu(self.bn_t_3(self.t_3(x_3)))
x_5 = self.relu(self.bn_m_5(self.m_5(x)))
x_5 = self.relu(self.bn_t_5(self.t_5(x_5)))
x = x_3 + x_5
x = self.relu(self.bn_last(self.last(x)))
return x
class audioEncoder(nn.Module):
def __init__(self):
super(audioEncoder, self).__init__()
self.block1 = Audio_Block(1, 32)
self.pool1 = nn.MaxPool3d(kernel_size = (1, 1, 3), stride = (1, 1, 2), padding = (0, 0, 1))
self.block2 = Audio_Block(32, 64)
self.pool2 = nn.MaxPool3d(kernel_size = (1, 1, 3), stride = (1, 1, 2), padding = (0, 0, 1))
self.block3 = Audio_Block(64, 128)
self.pool3 = nn.MaxPool3d(kernel_size = (1, 1, 3), stride = (1, 1, 1), padding = (0, 0, 1))
self.block4 = Audio_Block(128, 256)
self.__init_weight()
def forward(self, x):
x = self.block1(x)
x = self.pool1(x)
x = self.block2(x)
x = self.pool2(x)
x = self.block3(x)
x = self.pool3(x)
x = self.block4(x)
x = torch.mean(x, dim = 2, keepdim = True)
x = x.squeeze(2).transpose(1, 2)
return x
def __init_weight(self):
for m in self.modules():
if isinstance(m, nn.Conv2d):
torch.nn.init.kaiming_normal_(m.weight)
elif isinstance(m, nn.BatchNorm2d):
m.weight.data.fill_(1)
m.bias.data.zero_()
|