Views
No views yet
1import numpy as np
2import torch
3from pydub import AudioSegment
4
5from transformers import Wav2Vec2ForAudioFrameClassification, Wav2Vec2FeatureExtractor
6
7
8def _make_timegrid(sound_duration: float, total_len: int):
9 start_timegrid = np.linspace(0, sound_duration, total_len + 1)
10 dt = start_timegrid[1] - start_timegrid[0]
11 end_timegrid = start_timegrid + dt
12 return start_timegrid[:total_len], end_timegrid[:total_len]
13
14feature_extractor = Wav2Vec2FeatureExtractor(
15 feature_size=1,
16 sampling_rate=16_000,
17 padding_value=0.0,
18 do_normalize=True,
19 return_attention_mask=True,
20)
21model = Wav2Vec2ForAudioFrameClassification.from_pretrained("Ivydata/wav2vec2-large-speech-diarization-jp")
22filepath = "/path/to/file.wav"
23sound = AudioSegment.from_file(filepath)
24sound = sound.set_frame_rate(16_000)
25sound_duration = sound.duration_seconds
26
27feature = feature_extractor(np.array(sound.get_array_of_samples())).input_values[0]
28input_values = torch.tensor(feature, dtype=torch.float32).unsqueeze(0)
29
30with torch.no_grad():
31 logits = model(input_values).logits
32pred = logits.argmax(dim=-1).squeeze(0)
33start_timegrid, end_timegrid = _make_timegrid(sound_duration, len(pred))
34
35print("sec speaker_label")
36for p, start_time in zip(pred, start_timegrid):
37 print(f"{start_time:.4f} {p}")