RaphaelSchwinger commited on
Commit
2db0607
·
verified ·
1 Parent(s): 69ce718

Update preprocessing

Browse files
Files changed (1) hide show
  1. README.md +16 -13
README.md CHANGED
@@ -105,6 +105,15 @@ class PowerToDB(torch.nn.Module):
105
 
106
  return log_spec
107
 
 
 
 
 
 
 
 
 
 
108
 
109
  def preprocess(audio, sample_rate_of_audio):
110
  """
@@ -114,32 +123,26 @@ def preprocess(audio, sample_rate_of_audio):
114
  - Normalize the melscale spectrogram with mean: -4.268, std: 4.569 (from AudioSet)
115
 
116
  """
117
- powerToDB = PowerToDB()
118
- # Resample to 32kHz
119
- resample = torchaudio.transforms.Resample(
120
- orig_freq=sample_rate_of_audio, new_freq=32000
121
- )
122
- audio = resample(audio)
123
- spectrogram = torchaudio.transforms.Spectrogram(
124
- n_fft=2048, hop_length=256, power=2.0
125
- )(audio)
126
- melspec = torchaudio.transforms.MelScale(n_mels=256, n_stft=1025)(spectrogram)
127
  dbscale = powerToDB(melspec)
128
  normalized_dbscale = transforms.Normalize((-4.268,), (4.569,))(dbscale)
 
 
 
129
  return normalized_dbscale
130
 
131
  preprocessed_audio = preprocess(audio, sample_rate)
132
  print("Preprocessed_audio shape:", preprocessed_audio.shape)
133
 
134
-
135
- logits = model(preprocessed_audio.unsqueeze(0)).logits
136
  print("Logits shape: ", logits.shape)
137
 
138
  top5 = torch.topk(logits, 5)
139
  print("Top 5 logits:", top5.values)
140
  print("Top 5 predicted classes:")
141
  print([model.config.id2label[i] for i in top5.indices.squeeze().tolist()])
142
-
143
  ```
144
 
145
  ## Model Source
 
105
 
106
  return log_spec
107
 
108
+ # Initialize preprocessors
109
+ spectrogram_converter = torchaudio.transforms.Spectrogram(
110
+ n_fft=2048, hop_length=256, power=2.0
111
+ )
112
+ mel_converter = torchaudio.transforms.MelScale(
113
+ n_mels=256, n_stft=1025, sample_rate=32_000
114
+ )
115
+ powerToDB = PowerToDB(top_db=80)
116
+
117
 
118
  def preprocess(audio, sample_rate_of_audio):
119
  """
 
123
  - Normalize the melscale spectrogram with mean: -4.268, std: 4.569 (from AudioSet)
124
 
125
  """
126
+ spectrogram = spectrogram_converter(audio)
127
+ spectrogram = spectrogram.to(torch.float32)
128
+ melspec = mel_converter(spectrogram)
 
 
 
 
 
 
 
129
  dbscale = powerToDB(melspec)
130
  normalized_dbscale = transforms.Normalize((-4.268,), (4.569,))(dbscale)
131
+ # add batch dimension if needed
132
+ if normalized_dbscale.dim() == 3:
133
+ normalized_dbscale = normalized_dbscale.unsqueeze(0)
134
  return normalized_dbscale
135
 
136
  preprocessed_audio = preprocess(audio, sample_rate)
137
  print("Preprocessed_audio shape:", preprocessed_audio.shape)
138
 
139
+ logits = model(preprocessed_audio).logits
 
140
  print("Logits shape: ", logits.shape)
141
 
142
  top5 = torch.topk(logits, 5)
143
  print("Top 5 logits:", top5.values)
144
  print("Top 5 predicted classes:")
145
  print([model.config.id2label[i] for i in top5.indices.squeeze().tolist()])
 
146
  ```
147
 
148
  ## Model Source