DFAGWE/infinitetalk2
0
1from transformers import Wav2Vec2Config, Wav2Vec2Model2from transformers.modeling_outputs import BaseModelOutput3 4from src.audio_analysis.torch_utils import linear_interpolation5 6# the implementation of Wav2Vec2Model is borrowed from7# https://github.com/huggingface/transformers/blob/HEAD/src/transformers/models/wav2vec2/modeling_wav2vec2.py8# initialize our encoder with the pre-trained wav2vec 2.0 weights.9class Wav2Vec2Model(Wav2Vec2Model):10 def __init__(self, config: Wav2Vec2Config):11 super().__init__(config)12 13 def forward(14 self,15 input_values,16 seq_len,17 attention_mask=None,18 mask_time_indices=None,19 output_attentions=None,20 output_hidden_states=None,21 return_dict=None,22 ):23 self.config.output_attentions = True24 25 output_hidden_states = (26 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states27 )28 return_dict = return_dict if return_dict is not None else self.config.use_return_dict29 30 extract_features = self.feature_extractor(input_values)31 extract_features = extract_features.transpose(1, 2)32 extract_features = linear_interpolation(extract_features, seq_len=seq_len)33 34 if attention_mask is not None:35 # compute reduced attention_mask corresponding to feature vectors36 attention_mask = self._get_feature_vector_attention_mask(37 extract_features.shape[1], attention_mask, add_adapter=False38 )39 40 hidden_states, extract_features = self.feature_projection(extract_features)41 hidden_states = self._mask_hidden_states(42 hidden_states, mask_time_indices=mask_time_indices, attention_mask=attention_mask43 )44 45 encoder_outputs = self.encoder(46 hidden_states,47 attention_mask=attention_mask,48 output_attentions=output_attentions,49 output_hidden_states=output_hidden_states,50 return_dict=return_dict,51 )52 53 hidden_states = encoder_outputs[0]54 55 if self.adapter is not None:56 hidden_states = self.adapter(hidden_states)57 58 if not return_dict:59 return (hidden_states, ) + encoder_outputs[1:]60 return BaseModelOutput(61 last_hidden_state=hidden_states,62 hidden_states=encoder_outputs.hidden_states,63 attentions=encoder_outputs.attentions,64 )65 66 67 def feature_extract(68 self,69 input_values,70 seq_len,71 ):72 extract_features = self.feature_extractor(input_values)73 extract_features = extract_features.transpose(1, 2)74 extract_features = linear_interpolation(extract_features, seq_len=seq_len)75 76 return extract_features77 78 def encode(79 self,80 extract_features,81 attention_mask=None,82 mask_time_indices=None,83 output_attentions=None,84 output_hidden_states=None,85 return_dict=None,86 ):87 self.config.output_attentions = True88 89 output_hidden_states = (90 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states91 )92 return_dict = return_dict if return_dict is not None else self.config.use_return_dict93 94 if attention_mask is not None:95 # compute reduced attention_mask corresponding to feature vectors96 attention_mask = self._get_feature_vector_attention_mask(97 extract_features.shape[1], attention_mask, add_adapter=False98 )99 100 101 hidden_states, extract_features = self.feature_projection(extract_features)102 hidden_states = self._mask_hidden_states(103 hidden_states, mask_time_indices=mask_time_indices, attention_mask=attention_mask104 )105 106 encoder_outputs = self.encoder(107 hidden_states,108 attention_mask=attention_mask,109 output_attentions=output_attentions,110 output_hidden_states=output_hidden_states,111 return_dict=return_dict,112 )113 114 hidden_states = encoder_outputs[0]115 116 if self.adapter is not None:117 hidden_states = self.adapter(hidden_states)118 119 if not return_dict:120 return (hidden_states, ) + encoder_outputs[1:]121 return BaseModelOutput(122 last_hidden_state=hidden_states,123 hidden_states=encoder_outputs.hidden_states,124 attentions=encoder_outputs.attentions,125 )126 