Aluode/PerceptionLabPortable
0
1# coding=utf-82# Copyright 2022 The HuggingFace Inc. team.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15"""PyTorch Data2VecText model."""16 17import math18 19import torch20from torch import nn21 22from ...activations import ACT2FN23from ...modeling_layers import GradientCheckpointingLayer24from ...modeling_outputs import Wav2Vec2BaseModelOutput25from ...modeling_utils import PreTrainedModel26from ..wav2vec2.modeling_wav2vec2 import (27 Wav2Vec2Adapter,28 Wav2Vec2Encoder,29 Wav2Vec2FeatureEncoder,30 Wav2Vec2FeatureProjection,31 Wav2Vec2ForAudioFrameClassification,32 Wav2Vec2ForCTC,33 Wav2Vec2ForSequenceClassification,34 Wav2Vec2ForXVector,35 Wav2Vec2Model,36 Wav2Vec2PreTrainedModel,37 Wav2Vec2SamePadLayer,38)39from .configuration_data2vec_audio import Data2VecAudioConfig40 41 42class Data2VecAudioConvLayer(GradientCheckpointingLayer):43 def __init__(self, config, layer_id=0):44 super().__init__()45 self.in_conv_dim = config.conv_dim[layer_id - 1] if layer_id > 0 else 146 self.out_conv_dim = config.conv_dim[layer_id]47 48 self.conv = nn.Conv1d(49 self.in_conv_dim,50 self.out_conv_dim,51 kernel_size=config.conv_kernel[layer_id],52 stride=config.conv_stride[layer_id],53 bias=config.conv_bias,54 )55 self.layer_norm = nn.LayerNorm(self.out_conv_dim, elementwise_affine=True)56 self.activation = ACT2FN[config.feat_extract_activation]57 58 def forward(self, hidden_states):59 hidden_states = self.conv(hidden_states)60 61 hidden_states = hidden_states.transpose(-2, -1)62 hidden_states = self.layer_norm(hidden_states)63 hidden_states = hidden_states.transpose(-2, -1)64 65 hidden_states = self.activation(hidden_states)66 return hidden_states67 68 69class Data2VecAudioPadLayer(Wav2Vec2SamePadLayer):70 pass71 72 73class Data2VecAudioPositionalConvLayer(nn.Module):74 def __init__(self, config):75 super().__init__()76 self.conv = nn.Conv1d(77 config.hidden_size,78 config.hidden_size,79 kernel_size=config.conv_pos_kernel_size,80 padding=config.conv_pos_kernel_size // 2,81 groups=config.num_conv_pos_embedding_groups,82 )83 84 self.padding = Data2VecAudioPadLayer(config.conv_pos_kernel_size)85 self.activation = ACT2FN[config.feat_extract_activation]86 # no learnable parameters87 self.layer_norm = nn.LayerNorm(config.hidden_size, elementwise_affine=False)88 89 def forward(self, hidden_states):90 hidden_states = self.conv(hidden_states)91 hidden_states = self.padding(hidden_states)92 93 hidden_states = hidden_states.transpose(1, 2)94 hidden_states = self.layer_norm(hidden_states)95 hidden_states = hidden_states.transpose(1, 2)96 hidden_states = self.activation(hidden_states)97 return hidden_states98 99 100class Data2VecAudioPositionalConvEmbedding(nn.Module):101 def __init__(self, config):102 super().__init__()103 self.layers = nn.ModuleList(104 [Data2VecAudioPositionalConvLayer(config) for _ in range(config.num_conv_pos_embeddings)]105 )106 107 def forward(self, hidden_states):108 hidden_states = hidden_states.transpose(1, 2)109 for layer in self.layers:110 hidden_states = layer(hidden_states)111 hidden_states = hidden_states.transpose(1, 2)112 return hidden_states113 114 115class Data2VecAudioFeatureEncoder(Wav2Vec2FeatureEncoder):116 def __init__(self, config):117 nn.Module.__init__(self)118 self.conv_layers = nn.ModuleList(119 [Data2VecAudioConvLayer(config, layer_id=i) for i in range(config.num_feat_extract_layers)]120 )121 self.gradient_checkpointing = False122 self._requires_grad = True123 124 125class Data2VecAudioFeatureProjection(Wav2Vec2FeatureProjection):126 pass127 128 129class Data2VecAudioEncoder(Wav2Vec2Encoder):130 pass131 132 133class Data2VecAudioAdapter(Wav2Vec2Adapter):134 pass135 136 137class Data2VecAudioPreTrainedModel(PreTrainedModel, Wav2Vec2PreTrainedModel):138 config: Data2VecAudioConfig139 base_model_prefix = "data2vec_audio"140 main_input_name = "input_values"141 supports_gradient_checkpointing = True142 _supports_flash_attn = True143 _supports_sdpa = True144 _supports_flex_attn = True145 146 def _init_weights(self, module):147 """Initialize the weights"""148 if isinstance(module, Data2VecAudioFeatureProjection):149 k = math.sqrt(1 / module.projection.in_features)150 nn.init.uniform_(module.projection.weight, a=-k, b=k)151 nn.init.uniform_(module.projection.bias, a=-k, b=k)152 elif isinstance(module, Data2VecAudioPositionalConvLayer):153 nn.init.constant_(module.conv.bias, 0)154 elif isinstance(module, nn.Linear):155 module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)156 157 if module.bias is not None:158 module.bias.data.zero_()159 elif isinstance(module, (nn.LayerNorm, nn.GroupNorm)):160 if module.bias is not None:161 module.bias.data.zero_()162 if module.weight is not None:163 module.weight.data.fill_(1.0)164 elif isinstance(module, nn.Conv1d):165 nn.init.kaiming_normal_(module.weight)166 167 if module.bias is not None:168 k = math.sqrt(module.groups / (module.in_channels * module.kernel_size[0]))169 nn.init.uniform_(module.bias, a=-k, b=k)170 171 def _get_adapters(self):172 raise AttributeError("Not needed for Data2VecAudio")173 174 def init_adapter_layers(self):175 raise AttributeError("Not needed for Data2VecAudio")176 177 def load_adapter(self):178 raise AttributeError("Not needed for Data2VecAudio")179 180 181Data2VecAudioBaseModelOutput = Wav2Vec2BaseModelOutput182 183 184class Data2VecAudioModel(Data2VecAudioPreTrainedModel, Wav2Vec2Model):185 def __init__(self, config: Data2VecAudioConfig):186 Data2VecAudioPreTrainedModel.__init__(self, config)187 self.config = config188 self.feature_extractor = Data2VecAudioFeatureEncoder(config)189 self.feature_projection = Data2VecAudioFeatureProjection(config)190 191 # model only needs masking vector if mask prob is > 0.0192 if config.mask_time_prob > 0.0 or config.mask_feature_prob > 0.0:193 self.masked_spec_embed = nn.Parameter(torch.Tensor(config.hidden_size).uniform_())194 195 self.encoder = Data2VecAudioEncoder(config)196 197 self.adapter = Data2VecAudioAdapter(config) if config.add_adapter else None198 199 # Initialize weights and apply final processing200 self.post_init()201 202 def freeze_feature_extractor(self):203 raise AttributeError("Not needed for Data2VecAudio")204 205 def freeze_feature_encoder(self):206 """207 Calling this function will disable the gradient computation for the feature encoder so that its parameter will208 not be updated during training.209 """210 self.feature_extractor._freeze_parameters()211 212 def forward(self, **super_kwargs):213 return super().forward(**super_kwargs)214 215 216class Data2VecAudioForCTC(Data2VecAudioPreTrainedModel, Wav2Vec2ForCTC):217 def __init__(self, config):218 Data2VecAudioPreTrainedModel.__init__(self, config)219 220 self.data2vec_audio = Data2VecAudioModel(config)221 self.dropout = nn.Dropout(config.final_dropout)222 223 if config.vocab_size is None:224 raise ValueError(225 f"You are trying to instantiate {self.__class__} with a configuration that "226 "does not define the vocabulary size of the language model head. Please "227 "instantiate the model as follows: `Data2VecAudioForCTC.from_pretrained(..., vocab_size=vocab_size)`. "228 "or define `vocab_size` of your model's configuration."229 )230 output_hidden_size = (231 config.output_hidden_size if hasattr(config, "add_adapter") and config.add_adapter else config.hidden_size232 )233 self.lm_head = nn.Linear(output_hidden_size, config.vocab_size)234 235 # Initialize weights and apply final processing236 self.post_init()237 238 def freeze_base_model(self):239 raise AttributeError("Not needed for Data2VecAudio")240 241 def tie_weights(self):242 raise AttributeError("Not needed for Data2VecAudio")243 244 def forward(self, **super_kwargs):245 return super().forward(**super_kwargs)246 247 248class Data2VecAudioForSequenceClassification(Wav2Vec2ForSequenceClassification):249 pass250 251 252class Data2VecAudioForAudioFrameClassification(Wav2Vec2ForAudioFrameClassification):253 pass254 255 256class Data2VecAudioForXVector(Wav2Vec2ForXVector):257 pass258 259 260__all__ = [261 "Data2VecAudioForAudioFrameClassification",262 "Data2VecAudioForCTC",263 "Data2VecAudioForSequenceClassification",264 "Data2VecAudioForXVector",265 "Data2VecAudioModel",266 "Data2VecAudioPreTrainedModel",267]268 