shethjenil commited on
Commit
9754376
·
verified ·
1 Parent(s): 65f9021

Update modeling_conformer.py

Browse files
Files changed (1) hide show
  1. modeling_conformer.py +67 -121
modeling_conformer.py CHANGED
@@ -1,9 +1,8 @@
1
  from huggingface_hub import hf_hub_download
2
  from torch import nn
3
- from transformers import Wav2Vec2ConformerModel
4
  from safetensors.torch import load_file
5
  from torch_state_bridge import state_bridge
6
- import json
7
  import torch
8
  import torch.nn.functional as F
9
  import torchaudio
@@ -21,58 +20,28 @@ class Op(nn.Module):
21
  return self.func(x)
22
 
23
  class Wav2Vec2ConformerRNNT(Wav2Vec2ConformerModel):
24
- def __init__(self, config):
25
- self.language = config.languages[0]
26
- if len(config.languages) > 1:
27
- config.hidden_size = 1024
28
- config.num_hidden_layers = 24
29
- config.conv_depthwise_kernel_size = 9
30
- config.conv_stride = [2,2,2]
31
- config.conv_kernel = [3,3,3]
32
- config.conv_dim = [256,256,256]
33
- config.feat_extract_norm = "group"
34
- config.intermediate_size = 4096
35
- config.num_feat_extract_layers = len(config.conv_dim)
36
- config.lstm_layer = 2
37
-
38
- self.cache_length = None
39
- self.hop, self.preemph, self.eps, self.pad_to = 160, 0.97, 2**-24, 16
40
- self.denorm = (2 ** config.num_feat_extract_layers) * self.hop / config.sampling_rate
41
- self.scaler = config.hidden_size ** (1/2)
42
- super().__init__(config)
43
- self.eval()
44
 
45
  def init_weights(self):
46
  del self.encoder.pos_conv_embed
47
  config = self.config
 
48
  self.enc = nn.Linear(config.hidden_size, config.joint_hidden)
49
  self.pred = nn.Linear(config.pred_hidden, config.joint_hidden)
50
- self.joint = nn.Linear(config.joint_hidden, config.vocab_size // 22 + 1)
51
  self.embed = nn.Embedding(config.vocab_size+1, config.pred_hidden, padding_idx=config.vocab_size)
52
  self.lstm = nn.LSTM(config.pred_hidden, config.pred_hidden, config.lstm_layer, batch_first=True)
53
  self.act = nn.ReLU(inplace=True)
54
  self.spec = torchaudio.transforms.Spectrogram(n_fft=512, hop_length=160, win_length=400, center=False)
55
  self.mask_layer = Op(lambda self_obj,x : x.masked_fill(self_obj.cache_pad_mask.unsqueeze(1), 0),True)
56
- self.register_buffer(
57
- "mel_fb",
58
- torch.tensor(
59
- librosa.filters.mel(
60
- sr=self.config.sampling_rate,
61
- n_fft=512,
62
- n_mels=80
63
- )
64
- )
65
- )
66
-
67
  for idx,l in enumerate(self.feature_extractor.conv_layers):
68
- if len(self.config.languages) == 1 or idx == 0:
69
  l.conv = nn.Conv2d(l.conv.in_channels,l.conv.out_channels,l.conv.kernel_size[0],l.conv.stride,1)
70
  l.layer_norm = nn.Identity()
71
  else:
72
  l.conv = nn.Sequential(nn.Conv2d(l.conv.in_channels,l.conv.out_channels,l.conv.kernel_size[0],l.conv.stride,1,groups=l.conv.out_channels),nn.Conv2d(l.conv.in_channels,l.conv.out_channels, 1))
73
-
74
  self.feature_extractor.conv_layers.append(Op(lambda x : x.transpose(1, 2)))
75
- self.feature_projection.projection = nn.Linear(config.conv_dim[-1] * int(self.calc_length(torch.tensor(80.),repeat_num=self.config.num_feat_extract_layers)),config.hidden_size)
76
  self.feature_projection.layer_norm = Op(lambda x:x.permute(0, 2, 1, 3).flatten(2))
77
  for l in self.encoder.layers:
78
  l.conv_module.glu = nn.Sequential(l.conv_module.glu,self.mask_layer)
@@ -80,8 +49,11 @@ class Wav2Vec2ConformerRNNT(Wav2Vec2ConformerModel):
80
  l.conv_module.pointwise_conv2.bias = nn.Parameter(torch.empty(l.conv_module.pointwise_conv2.out_channels))
81
  l.conv_module.depthwise_conv.bias = nn.Parameter(torch.empty(l.conv_module.depthwise_conv.out_channels))
82
  self.encoder.layer_norm = nn.Identity()
83
- if len(self.config.languages) > 1:
84
- self.lang_joint_net = nn.ModuleDict({l: nn.Linear(config.joint_hidden, config.vocab_size // 22 + 1) for l in config.languages})
 
 
 
85
  return super().init_weights()
86
 
87
  def _mask_hidden_states(self, hidden_states, mask_time_indices = None, attention_mask = None):
@@ -97,7 +69,7 @@ class Wav2Vec2ConformerRNNT(Wav2Vec2ConformerModel):
97
 
98
  def preprocessing(self, x):
99
  x, l = x
100
- l = (l // self.hop + 1).long()
101
  x = torch.cat((x[:, :1], x[:, 1:] - self.preemph * x[:, :-1]), 1)
102
  x = (self.mel_fb @ self.spec(x) + self.eps).log()
103
  T = x.size(-1)
@@ -113,72 +85,23 @@ class Wav2Vec2ConformerRNNT(Wav2Vec2ConformerModel):
113
  def forward(self, input_values):
114
  return self._greedy_decode(super().forward(self.preprocessing(input_values)).last_hidden_state)
115
 
116
- def load_state_dict(self, state_dict, strict=True, assign=False):
117
- state_dict.pop('ctc_decoder.decoder_layers.0.bias', None)
118
- state_dict.pop('ctc_decoder.decoder_layers.0.weight', None)
119
-
120
- state_dict['preprocessor.featurizer.fb'] = state_dict['preprocessor.featurizer.fb'].squeeze(0)
121
- changes = """
122
- preprocessor.featurizer.fb,mel_fb
123
- preprocessor.featurizer.window,spec.window
124
- norm_feed_forward1,ffn1_layer_norm
125
- norm_feed_forward2,ffn2_layer_norm
126
- feed_forward1.linear1,ffn1.intermediate_dense
127
- feed_forward1.linear2,ffn1.output_dense
128
- feed_forward2.linear1,ffn2.intermediate_dense
129
- feed_forward2.linear2,ffn2.output_dense
130
- norm_self_att,self_attn_layer_norm
131
- norm_out,final_layer_norm
132
- norm_conv,conv_module.layer_norm
133
- .conv.,.conv_module.
134
- decoder.prediction.dec_rnn.lstm,lstm
135
- decoder.prediction.embed,embed
136
- joint.enc,enc
137
- joint.pred,pred
138
- joint.joint_net.2,lang_joint_net
139
- encoder.pre_encode.conv_module.0,feature_extractor.conv_layers.0.conv
140
- encoder.pre_encode.out,feature_projection.projection
141
- """
142
- if len(self.config.languages) == 1:
143
- changes += f"""lang_joint_net.{self.language},joint
144
- encoder.pre_encode.conv_module.{{n}},feature_extractor.conv_layers.{{(n/2)}}.conv"""
145
- else:
146
- state_dict["joint.weight"] = self.joint.weight.clone()
147
- state_dict["joint.bias"] = self.joint.bias.clone()
148
- changes += """encoder.pre_encode.conv_module.{n},encoder.pre_encode.conv_module.{(n-2)}
149
- encoder.pre_encode.conv_module.{n},feature_extractor.conv_layers.{(n//3+1)}.conv.{(n%3)}
150
- """
151
- # replicate many changes for complex maths
152
- state_dict = state_bridge(state_dict, changes)
153
- if len(self.config.languages) == 1:
154
- state_dict = {k: v for k, v in state_dict.items() if "lang_joint_net" not in k}
155
- return super().load_state_dict(state_dict, strict, assign)
156
-
157
- @torch.jit.export
158
  def _greedy_decode(self, enc_out: torch.Tensor):
159
-
160
  B, T, _ = enc_out.size()
161
  device = enc_out.device
162
-
163
  enc_proj = self.enc(enc_out)
164
-
165
  max_symbols = self.config.max_symbols_per_step
166
  max_len = T * max_symbols
167
-
168
  token_buffer = torch.full(
169
  (B, max_len),
170
  -1,
171
  dtype=torch.long,
172
  device=device
173
  )
174
-
175
  start_buffer = torch.zeros(
176
  (B, max_len),
177
  device=device
178
  )
179
-
180
  lengths = torch.zeros(B, dtype=torch.long, device=device)
181
-
182
  last = torch.full(
183
  (B, 1),
184
  self.config.blank_id,
@@ -230,37 +153,60 @@ encoder.pre_encode.conv_module.{n},feature_extractor.conv_layers.{(n//3+1)}.conv
230
 
231
  return tokens, starts
232
 
233
- @classmethod
234
- def from_pretrained(
235
- cls,
236
- pretrained_model_name_or_path,
237
- config=None,
238
- language=None,
239
- use_quantization=False):
240
-
241
- if config is None:
242
- raise ValueError("config must be provided")
243
-
244
- if language:
245
- config.languages = [language]
246
-
247
- vocab_file = hf_hub_download(
248
- pretrained_model_name_or_path,
249
- "vocab.json"
250
- )
251
-
252
- vocab_json = json.load(open(vocab_file))
253
- config.vocab = ['<unk>'] + vocab_json['small'][language]
254
-
255
- model = cls(config)
256
-
257
- weight_file = hf_hub_download(
258
- pretrained_model_name_or_path,
259
- f"{language or 'all'}.safetensors"
260
- )
261
 
262
- model.load_state_dict(load_file(weight_file))
263
- if use_quantization:
264
- model = torch.quantization.quantize_dynamic(model)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
265
 
266
- return model
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  from huggingface_hub import hf_hub_download
2
  from torch import nn
3
+ from transformers import Wav2Vec2ConformerModel , Wav2Vec2CTCTokenizer
4
  from safetensors.torch import load_file
5
  from torch_state_bridge import state_bridge
 
6
  import torch
7
  import torch.nn.functional as F
8
  import torchaudio
 
20
  return self.func(x)
21
 
22
  class Wav2Vec2ConformerRNNT(Wav2Vec2ConformerModel):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
23
 
24
  def init_weights(self):
25
  del self.encoder.pos_conv_embed
26
  config = self.config
27
+ self.cache_length = None
28
  self.enc = nn.Linear(config.hidden_size, config.joint_hidden)
29
  self.pred = nn.Linear(config.pred_hidden, config.joint_hidden)
30
+ self.joint = nn.Linear(config.joint_hidden, config.vocab_size // len(config.languages) + 1)
31
  self.embed = nn.Embedding(config.vocab_size+1, config.pred_hidden, padding_idx=config.vocab_size)
32
  self.lstm = nn.LSTM(config.pred_hidden, config.pred_hidden, config.lstm_layer, batch_first=True)
33
  self.act = nn.ReLU(inplace=True)
34
  self.spec = torchaudio.transforms.Spectrogram(n_fft=512, hop_length=160, win_length=400, center=False)
35
  self.mask_layer = Op(lambda self_obj,x : x.masked_fill(self_obj.cache_pad_mask.unsqueeze(1), 0),True)
36
+ self.register_buffer("mel_fb",torch.tensor(librosa.filters.mel(sr=config.sampling_rate,n_fft=512,n_mels=80)))
 
 
 
 
 
 
 
 
 
 
37
  for idx,l in enumerate(self.feature_extractor.conv_layers):
38
+ if not(config.multilingual) or idx == 0:
39
  l.conv = nn.Conv2d(l.conv.in_channels,l.conv.out_channels,l.conv.kernel_size[0],l.conv.stride,1)
40
  l.layer_norm = nn.Identity()
41
  else:
42
  l.conv = nn.Sequential(nn.Conv2d(l.conv.in_channels,l.conv.out_channels,l.conv.kernel_size[0],l.conv.stride,1,groups=l.conv.out_channels),nn.Conv2d(l.conv.in_channels,l.conv.out_channels, 1))
 
43
  self.feature_extractor.conv_layers.append(Op(lambda x : x.transpose(1, 2)))
44
+ self.feature_projection.projection = nn.Linear(config.conv_dim[-1] * int(self.calc_length(torch.tensor(80.),repeat_num=config.num_feat_extract_layers)),config.hidden_size)
45
  self.feature_projection.layer_norm = Op(lambda x:x.permute(0, 2, 1, 3).flatten(2))
46
  for l in self.encoder.layers:
47
  l.conv_module.glu = nn.Sequential(l.conv_module.glu,self.mask_layer)
 
49
  l.conv_module.pointwise_conv2.bias = nn.Parameter(torch.empty(l.conv_module.pointwise_conv2.out_channels))
50
  l.conv_module.depthwise_conv.bias = nn.Parameter(torch.empty(l.conv_module.depthwise_conv.out_channels))
51
  self.encoder.layer_norm = nn.Identity()
52
+ if config.multilingual:
53
+ self.lang_joint_net = nn.ModuleDict({l: nn.Linear(config.joint_hidden, config.vocab_size // len(config.languages) + 1) for l in config.languages})
54
+ self.preemph, self.eps, self.pad_to = 0.97, 2**-24, 16
55
+ self.denorm = (2 ** config.num_feat_extract_layers) * self.spec.hop_length / config.sampling_rate
56
+ self.scaler = config.hidden_size ** (1/2)
57
  return super().init_weights()
58
 
59
  def _mask_hidden_states(self, hidden_states, mask_time_indices = None, attention_mask = None):
 
69
 
70
  def preprocessing(self, x):
71
  x, l = x
72
+ l = (l // self.spec.hop_length + 1).long()
73
  x = torch.cat((x[:, :1], x[:, 1:] - self.preemph * x[:, :-1]), 1)
74
  x = (self.mel_fb @ self.spec(x) + self.eps).log()
75
  T = x.size(-1)
 
85
  def forward(self, input_values):
86
  return self._greedy_decode(super().forward(self.preprocessing(input_values)).last_hidden_state)
87
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
88
  def _greedy_decode(self, enc_out: torch.Tensor):
 
89
  B, T, _ = enc_out.size()
90
  device = enc_out.device
 
91
  enc_proj = self.enc(enc_out)
 
92
  max_symbols = self.config.max_symbols_per_step
93
  max_len = T * max_symbols
 
94
  token_buffer = torch.full(
95
  (B, max_len),
96
  -1,
97
  dtype=torch.long,
98
  device=device
99
  )
 
100
  start_buffer = torch.zeros(
101
  (B, max_len),
102
  device=device
103
  )
 
104
  lengths = torch.zeros(B, dtype=torch.long, device=device)
 
105
  last = torch.full(
106
  (B, 1),
107
  self.config.blank_id,
 
153
 
154
  return tokens, starts
155
 
156
+ def change_language(self,language):
157
+ self.joint.load_state_dict(self.lang_joint_net[language].state_dict())
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
158
 
159
+ @classmethod
160
+ def from_pretrained(cls, pretrained_model_name_or_path, *model_args, config = None, cache_dir = None, ignore_mismatched_sizes = False, force_download = False, local_files_only = False, token = None, revision = "main", use_safetensors = None, weights_only = True, **kwargs):
161
+ config.language = kwargs.pop("language",None)
162
+ config.multilingual = not(config.language)
163
+ if config.multilingual:
164
+ config.hidden_size = 1024
165
+ config.num_hidden_layers = 24
166
+ config.conv_depthwise_kernel_size = 9
167
+ config.conv_stride = [2,2,2]
168
+ config.conv_kernel = [3,3,3]
169
+ config.conv_dim = [256,256,256]
170
+ config.feat_extract_norm = "group"
171
+ config.intermediate_size = config.hidden_size * 4
172
+ config.num_feat_extract_layers = len(config.conv_dim)
173
+ config.lstm_layer = 2
174
+ kwargs['state_dict'] = load_file(hf_hub_download(pretrained_model_name_or_path,f"{config.language or 'all'}.safetensors"))
175
+ return super().from_pretrained(None, *model_args, config=config, cache_dir=cache_dir, ignore_mismatched_sizes=ignore_mismatched_sizes, force_download=force_download, local_files_only=local_files_only, token=token, revision=revision, use_safetensors=use_safetensors, weights_only=weights_only, **kwargs)
176
 
177
+ @staticmethod
178
+ def _load_pretrained_model(model, state_dict, checkpoint_files, load_config):
179
+ changes = """
180
+ preprocessor.featurizer.fb,mel_fb
181
+ preprocessor.featurizer.window,spec.window
182
+ norm_feed_forward1,ffn1_layer_norm
183
+ norm_feed_forward2,ffn2_layer_norm
184
+ feed_forward1.linear1,ffn1.intermediate_dense
185
+ feed_forward1.linear2,ffn1.output_dense
186
+ feed_forward2.linear1,ffn2.intermediate_dense
187
+ feed_forward2.linear2,ffn2.output_dense
188
+ norm_self_att,self_attn_layer_norm
189
+ norm_out,final_layer_norm
190
+ norm_conv,conv_module.layer_norm
191
+ .conv.,.conv_module.
192
+ decoder.prediction.dec_rnn.lstm,lstm
193
+ decoder.prediction.embed,embed
194
+ joint.enc,enc
195
+ joint.pred,pred
196
+ joint.joint_net.2,lang_joint_net
197
+ encoder.pre_encode.conv_module.0,feature_extractor.conv_layers.0.conv
198
+ encoder.pre_encode.out,feature_projection.projection
199
+ """
200
+ if not model.config.multilingual:
201
+ changes += "encoder.pre_encode.conv_module.{n},feature_extractor.conv_layers.{(n/2)}.conv"
202
+ changes += f"lang_joint_net.{model.config.language},joint"
203
+ else:
204
+ changes += "encoder.pre_encode.conv_module.{n},encoder.pre_encode.conv_module.{(n-2)}"
205
+ changes += "encoder.pre_encode.conv_module.{n},feature_extractor.conv_layers.{(n//3+1)}.conv.{(n%3)}"
206
+ state_dict = state_bridge(state_dict, changes)
207
+ if not model.config.multilingual:
208
+ state_dict = {k: v for k, v in state_dict.items() if "lang_joint_net" not in k}
209
+ state_dict['mel_fb'] = state_dict['mel_fb'].squeeze(0)
210
+ state_dict.pop('ctc_decoder.decoder_layers.0.bias', None)
211
+ state_dict.pop('ctc_decoder.decoder_layers.0.weight', None)
212
+ return super()._load_pretrained_model(model, state_dict, checkpoint_files, load_config)