File size: 2,668 Bytes
2566adf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Frozen Whisper encoder wrapper (wav -> encoder embeddings)."""

from __future__ import annotations

from typing import Optional

import torch
import torch.nn as nn
import torchaudio
from transformers import WhisperFeatureExtractor, WhisperModel

WHISPER_MEL_FRAMES = 3000        # 3000 frames at 16000 Hz


class WhisperEncoder():

    def __init__(
        self,
        device: Optional[str] = None,
    ) -> None:
        self.fe = WhisperFeatureExtractor.from_pretrained("openai/whisper-base")
        self.model = WhisperModel.from_pretrained("openai/whisper-base")
        self.model = self.model.to(device or ("cuda" if torch.cuda.is_available() else "cpu"))

    def encode(
        self,
        wav: torch.Tensor,
        sr: int,
    ) -> torch.Tensor:
        wav = torchaudio.functional.resample(wav, orig_freq=sr, new_freq=self.fe.sampling_rate) if sr != self.fe.sampling_rate else wav
        wav_np = wav.cpu().detach().numpy().astype("float32", copy=False)

        inputs = self.fe(
            wav_np,
            sampling_rate=self.fe.sampling_rate,
            return_tensors="pt",
            padding=False,
            truncation=False,
            return_attention_mask=True,
        )

        input_features = inputs.input_features
        num_frames = input_features.shape[-1]
        if num_frames < WHISPER_MEL_FRAMES:
            pad = WHISPER_MEL_FRAMES - num_frames
            input_features = torch.nn.functional.pad(input_features, (0, pad))
        else:
            input_features = input_features[..., :WHISPER_MEL_FRAMES]

        input_features = input_features.to(wav.device)
        if self.model.device != wav.device:
            self.model = self.model.to(wav.device)
        attention_mask = inputs.attention_mask.to(wav.device) if inputs.attention_mask is not None else None

        encoder_out = self.model.encoder(input_features).last_hidden_state

        if attention_mask is not None:
            valid_mel_frames = attention_mask.sum(dim=1)
            valid_enc_frames = (valid_mel_frames + 1) // 2
            max_valid_enc_frames = min(int(valid_enc_frames.max().item()), encoder_out.shape[1])
            encoder_out = encoder_out[:, :max_valid_enc_frames, :]
            valid_len = min(int(valid_enc_frames[0].item()), max_valid_enc_frames)
            if valid_len < max_valid_enc_frames:
                encoder_out[0, valid_len:, :] = 0

        return encoder_out


if __name__ == "__main__":
    torch.manual_seed(0)
    audio = torch.randn(1, 24000 * 25).float().to("cuda")
    encoder = WhisperEncoder()
    whisper_encoder_out = encoder.encode(audio, sr=24000)
    print(whisper_encoder_out.shape)