Views
No views yet
1
2
3import torch
4import torch.nn.functional as F
5import soundfile as sf
6from fairseq import checkpoint_utils
7
8from transformers import (
9 Wav2Vec2FeatureExtractor,
10 Wav2Vec2ForPreTraining,
11 Wav2Vec2Model,
12)
13from transformers.models.wav2vec2.modeling_wav2vec2 import _compute_mask_indices
14
15model_path=""
16wav_path=""
17mask_prob=0.0
18mask_length=10
19
20feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(model_path)
21model = Wav2Vec2Model.from_pretrained(model_path)
22
23# for pretrain: Wav2Vec2ForPreTraining
24# model = Wav2Vec2ForPreTraining.from_pretrained(model_path)
25
26model = model.to(device)
27model = model.half()
28model.eval()
29
30wav, sr = sf.read(wav_path)
31input_values = feature_extractor(wav, return_tensors="pt").input_values
32input_values = input_values.half()
33input_values = input_values.to(device)
34
35# for Wav2Vec2ForPreTraining
36# batch_size, raw_sequence_length = input_values.shape
37# sequence_length = model._get_feat_extract_output_lengths(raw_sequence_length)
38# mask_time_indices = _compute_mask_indices((batch_size, sequence_length), mask_prob=0.0, mask_length=2)
39# mask_time_indices = torch.tensor(mask_time_indices, device=input_values.device, dtype=torch.long)
40
41with torch.no_grad():
42 outputs = model(input_values)
43 last_hidden_state = outputs.last_hidden_state
44
45 # for Wav2Vec2ForPreTraining
46 # outputs = model(input_values, mask_time_indices=mask_time_indices, output_hidden_states=True)
47 # last_hidden_state = outputs.hidden_states[-1]
48