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_()