CoolFace
Apppublic

DFAGWE/infinitetalk2

sourceHugging Faceapache-2.0updated 8mo agoView on Hugging Face
0likes
wav2vec2.py126 linesDownload Raw Back to audio_analysis
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