| 18 | |
| 19 | |
| 20 | class AudioEncoder(nn.Module): |
| 21 | def __init__(self, path): |
| 22 | super().__init__() |
| 23 | self.model = torch.jit.load(path) |
| 24 | self.register_buffer('hidden', torch.zeros(2, 1, 256)) |
| 25 | |
| 26 | def forward(self, audio): |
| 27 | self.reset() |
| 28 | x = create_windowed_sequence(audio, 3200, cutting_stride=640, pad_samples=3200-640, cut_dim=1) |
| 29 | embs = [] |
| 30 | for i in range(x.shape[1]): |
| 31 | emb, _, self.hidden = self.model(x[:, i], torch.LongTensor([3200]), init_state=self.hidden) |
| 32 | embs.append(emb) |
| 33 | return torch.vstack(embs) |
| 34 | |
| 35 | def reset(self): |
| 36 | self.hidden = torch.zeros(2, 1, 256).to(self.hidden.device) |
| 37 | |
| 38 | |
| 39 | def get_audio_emb(audio_path, checkpoint, device): |