RASMUS/Finnish-ASR-Canary-v2
01.2k
1# Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7# http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14# Example Run Command: python make_supdata.py --ssl_model_ckpt_path <PATH TO CKPT> --manifest_path <PATH TO MANIFEST>15 16import argparse17import json18import os19import time20from multiprocessing import Pool21from pathlib import Path22 23import hydra.utils24import librosa25import numpy as np26import torch27from omegaconf import open_dict28from tqdm import tqdm29 30from nemo.collections.asr.parts.preprocessing.segment import AudioSegment31from nemo.collections.tts.models import ssl_tts32from nemo.collections.tts.parts.utils.tts_dataset_utils import get_base_dir33from nemo.core.classes import Dataset34from nemo.utils import logging35 36 37class AudioDataset(Dataset):38 def __init__(39 self,40 manifest_paths,41 min_duration=0.5,42 max_duration=16.0,43 pad_multiple=1024,44 sample_rate=22050,45 sup_data_dir=None,46 ):47 self.data = []48 for manifest_path in manifest_paths:49 with open(manifest_path, "r") as f:50 for line in f:51 record = json.loads(line)52 if record['duration'] < min_duration or record['duration'] > max_duration:53 continue54 self.data.append(json.loads(line))55 56 self.base_data_dir = get_base_dir([item["audio_filepath"] for item in self.data])57 if sup_data_dir is not None:58 self.sup_data_dir = sup_data_dir59 else:60 self.sup_data_dir = os.path.join(self.base_data_dir, "sup_data")61 if not os.path.exists(self.sup_data_dir):62 os.makedirs(self.sup_data_dir)63 64 self.pad_multiple = pad_multiple65 self.sample_rate = sample_rate66 67 def __len__(self):68 return len(self.data)69 70 def _get_wav_from_filepath(self, audio_filepath):71 features = AudioSegment.segment_from_file(72 audio_filepath, target_sr=self.sample_rate, n_segments=-1, trim=False,73 )74 audio_samples = features.samples75 audio, audio_length = torch.tensor(audio_samples), torch.tensor(audio_samples.shape[0]).long()76 77 # pad audio to a multiple of self.pad_multiple78 if audio.shape[0] % self.pad_multiple != 0:79 audio = torch.cat(80 [audio, torch.zeros(self.pad_multiple - audio.shape[0] % self.pad_multiple, dtype=torch.float)]81 )82 audio_length = torch.tensor(audio.shape[0]).long()83 84 return audio, audio_length85 86 def pad_collate_fn(self, batch):87 final_batch = {}88 for row in batch:89 for key in row:90 if key not in final_batch:91 final_batch[key] = []92 final_batch[key].append(row[key])93 94 max_audio_len = max([_audio_len.item() for _audio_len in final_batch["audio_len"]])95 96 audios_padded = []97 for audio in final_batch["audio"]:98 audio_padded = torch.nn.functional.pad(audio, (0, max_audio_len - audio.size(0)), value=0)99 audios_padded.append(audio_padded)100 101 final_batch["audio"] = audios_padded102 for key in final_batch:103 if key not in ["rel_audio_path_as_text_id", "wav_path"]:104 final_batch[key] = torch.stack(final_batch[key])105 106 return final_batch107 108 def __getitem__(self, index):109 sample = self.data[index]110 rel_audio_path = Path(sample["audio_filepath"]).relative_to(self.base_data_dir).with_suffix("")111 rel_audio_path_as_text_id = str(rel_audio_path).replace("/", "_")112 speaker = torch.tensor(sample["speaker"]).long()113 114 audio, audio_length = self._get_wav_from_filepath(sample["audio_filepath"])115 116 return {117 "audio": audio,118 "audio_len": audio_length,119 "rel_audio_path_as_text_id": rel_audio_path_as_text_id,120 "wav_path": sample["audio_filepath"],121 "speaker": speaker,122 }123 124 125def segment_wav(wav, segment_length, segment_hop_size, min_segment_length):126 if len(wav) < segment_length:127 pad = torch.zeros(segment_length - len(wav))128 segment = torch.cat([wav, pad])129 return [segment]130 else:131 si = 0132 segments = []133 while si < len(wav) - min_segment_length:134 segment = wav[si : si + segment_length]135 if len(segment) < segment_length:136 pad = torch.zeros(segment_length - len(segment))137 segment = torch.cat([segment, pad])138 segments.append(segment)139 si += segment_hop_size140 return segments141 142 143def segment_batch(batch, segment_length=44100, segment_hop_size=22050, min_segment_length=22050):144 all_segments = []145 segment_indices = []146 si = 0147 for bidx in range(len(batch['audio'])):148 audio = batch['audio'][bidx]149 audio_length = batch['audio_len'][bidx]150 audio_actual = audio[:audio_length]151 audio_segments = segment_wav(audio_actual, segment_length, segment_hop_size, min_segment_length)152 all_segments += audio_segments153 segment_indices.append((si, si + len(audio_segments) - 1))154 si += len(audio_segments)155 156 return torch.stack(all_segments), segment_indices157 158 159def get_mel_spectrogram(fb, wav, stft_params):160 EPSILON = 1e-9161 window_fn = torch.hann_window162 163 spec = torch.stft(164 input=wav,165 n_fft=stft_params['n_fft'], # 1024166 hop_length=stft_params['hop_length'], # 256167 win_length=stft_params['win_length'], # 1024168 window=window_fn(stft_params['win_length'], periodic=False).to(torch.float).to('cuda') if window_fn else None,169 return_complex=True,170 center=True,171 )172 173 if spec.dtype in [torch.cfloat, torch.cdouble]:174 spec = torch.view_as_real(spec)175 spec = torch.sqrt(spec.pow(2).sum(-1) + EPSILON)176 177 mel = torch.matmul(fb.to(spec.dtype), spec)178 log_mel = torch.log(torch.clamp(mel, min=torch.finfo(mel.dtype).tiny))179 180 return log_mel181 182 183def load_wav(wav_path, sample_rate=22050, pad_multiple=1024):184 wav = AudioSegment.segment_from_file(wav_path, target_sr=sample_rate, n_segments=-1, trim=False,).samples185 186 if wav.shape[0] % pad_multiple != 0:187 wav = np.concatenate([wav, np.zeros(pad_multiple - wav.shape[0] % pad_multiple)])188 wav = wav[:-1]189 190 return wav191 192 193def save_pitch_contour(record):194 wav_path = record['wav_path']195 wav_text_id = record['wav_id']196 sup_data_dir = record['sup_data_dir']197 stft_params = record['stft_params']198 wav = load_wav(wav_path, stft_params['sample_rate'], stft_params['pad_multiple'])199 pitch_contour_fn = f"pitch_contour_{wav_text_id}.pt"200 pitch_contour_fp = os.path.join(sup_data_dir, pitch_contour_fn)201 202 f0, _, _ = librosa.pyin(203 wav,204 fmin=librosa.note_to_hz('C2'),205 fmax=stft_params['yin_fmax'],206 frame_length=stft_params['win_length'],207 hop_length=stft_params['hop_length'],208 sr=stft_params['sample_rate'],209 center=True,210 fill_na=0.0,211 )212 213 pitch_contour = torch.tensor(f0, dtype=torch.float32)214 torch.save(pitch_contour, pitch_contour_fp)215 logging.info("saved {}".format(pitch_contour_fp))216 217 return pitch_contour218 219 220def compute_pitch_stats(records):221 def _is_valid_pitch(pitch_mean, pitch_std):222 c1 = pitch_mean > 0 and pitch_mean < 1000223 c2 = pitch_std > 0 and pitch_std < 1000224 return c1 and c2225 226 speaker_wise_pitch_contours = {}227 for item in records:228 wav_id = item['wav_id']229 speaker = item['speaker']230 sup_data_dir = item['sup_data_dir']231 pitch_contour_fn = f"pitch_contour_{wav_id}.pt"232 pitch_contour_fp = os.path.join(sup_data_dir, pitch_contour_fn)233 if speaker not in speaker_wise_pitch_contours:234 speaker_wise_pitch_contours[speaker] = []235 speaker_wise_pitch_contours[speaker].append(pitch_contour_fp)236 237 speaker_pitch_stats = {}238 for speaker in speaker_wise_pitch_contours:239 non_zero_pc = []240 for pitch_contour_fp in speaker_wise_pitch_contours[speaker][:50]:241 pitch_contour = torch.load(pitch_contour_fp)242 pitch_contour_nonzero = pitch_contour[pitch_contour != 0]243 if len(pitch_contour_nonzero) > 0:244 non_zero_pc.append(pitch_contour_nonzero)245 246 if len(non_zero_pc) > 0:247 non_zero_pc = torch.cat(non_zero_pc)248 pitch_mean = non_zero_pc.mean().item()249 pitch_std = non_zero_pc.std().item()250 valid = True251 252 if not _is_valid_pitch(pitch_mean, pitch_std):253 logging.warning("invalid pitch: {}".format(speaker))254 pitch_mean = 212.0255 pitch_std = 70.0256 valid = "False"257 else:258 logging.warning("could not find pitch contour for speaker {}".format(speaker))259 valid = "False"260 pitch_mean = 212.0261 pitch_std = 70.0262 263 speaker_pitch_stats[speaker] = {"pitch_mean": pitch_mean, "pitch_std": pitch_std, "valid": valid}264 265 with open(os.path.join(sup_data_dir, "speaker_pitch_stats.json"), "w") as f:266 json.dump(speaker_pitch_stats, f)267 268 269def main():270 parser = argparse.ArgumentParser(description='Evaluate the model')271 parser.add_argument(272 '--ssl_model_ckpt_path', type=str, required=True,273 )274 parser.add_argument('--manifest_paths', type=str, required=True)275 parser.add_argument('--sup_data_dir', type=str, default=None)276 parser.add_argument('--batch_size', type=int, default=32)277 parser.add_argument('--ssl_content_emb_type', type=str, default="embedding_and_probs")278 parser.add_argument('--use_unique_tokens', type=int, default=1)279 parser.add_argument('--num_workers', type=int, default=8)280 parser.add_argument('--pool_workers', type=int, default=30)281 parser.add_argument('--compute_pitch_contours', type=int, default=1)282 parser.add_argument('--num_pitch_per_speaker', type=int, default=None) # saves time.283 parser.add_argument('--sample_rate', type=int, default=22050)284 parser.add_argument('--pad_multiple', type=int, default=1024)285 parser.add_argument('--ssl_downsampling_factor', type=int, default=4)286 parser.add_argument('--stft_n_fft', type=int, default=1024)287 parser.add_argument('--stft_hop_length', type=int, default=256)288 parser.add_argument('--stft_win_length', type=int, default=1024)289 parser.add_argument('--stft_n_mel', type=int, default=80)290 parser.add_argument('--stft_fmin', type=int, default=0)291 parser.add_argument('--stft_fmax', type=int, default=8000)292 parser.add_argument('--yin_fmax', type=int, default=500)293 parser.add_argument('--segment_length', type=int, default=44100)294 parser.add_argument('--segment_hop_size', type=int, default=22050)295 parser.add_argument('--min_segment_length', type=int, default=22050)296 297 args = parser.parse_args()298 299 device = "cuda:0" if torch.cuda.is_available() else "cpu"300 301 manifest_paths = args.manifest_paths.split(",")302 ssl_model_ckpt_path = args.ssl_model_ckpt_path303 304 dataset = AudioDataset(305 manifest_paths, pad_multiple=args.pad_multiple, sample_rate=args.sample_rate, sup_data_dir=args.sup_data_dir306 )307 dataloader = torch.utils.data.DataLoader(308 dataset,309 batch_size=args.batch_size,310 shuffle=False,311 collate_fn=dataset.pad_collate_fn,312 num_workers=args.num_workers,313 )314 315 ssl_model = ssl_tts.SSLDisentangler.load_from_checkpoint(ssl_model_ckpt_path, strict=False)316 with open_dict(ssl_model.cfg):317 ssl_model.cfg.preprocessor.exact_pad = True318 ssl_model.preprocessor = hydra.utils.instantiate(ssl_model.cfg.preprocessor)319 ssl_model.preprocessor_disentangler = ssl_model.preprocessor320 ssl_model.eval()321 ssl_model.to(device)322 323 sample_rate = args.sample_rate324 stft_params = {325 "n_fft": args.stft_n_fft,326 "hop_length": args.stft_hop_length,327 "win_length": args.stft_win_length,328 "n_mel": args.stft_n_mel,329 "sample_rate": sample_rate,330 "pad_multiple": args.pad_multiple,331 "fmin": args.stft_fmin,332 "fmax": args.stft_fmax,333 "yin_fmax": args.yin_fmax,334 }335 336 fb = (337 torch.tensor(338 librosa.filters.mel(339 sr=sample_rate,340 n_fft=stft_params['n_fft'],341 n_mels=stft_params['n_mel'],342 fmin=stft_params['fmin'],343 fmax=stft_params['fmax'],344 ),345 dtype=torch.float,346 )347 .unsqueeze(0)348 .to(device)349 )350 351 st = time.time()352 bidx = 0353 wav_and_id_list = []354 355 for batch in tqdm(dataloader):356 bidx += 1357 with torch.no_grad():358 (359 _,360 _,361 batch_content_embedding,362 batch_content_log_probs,363 batch_encoded_len,364 ) = ssl_model.forward_for_export(365 input_signal=batch['audio'].to(device),366 input_signal_length=batch['audio_len'].to(device),367 normalize_content=True,368 )369 370 batch_mel_specs = get_mel_spectrogram(fb, batch['audio'][:, :-1].to(device), stft_params)371 audio_segmented, segment_indices = segment_batch(372 batch, args.segment_length, args.segment_hop_size, args.min_segment_length373 )374 audio_seg_len = torch.tensor([len(segment) for segment in audio_segmented]).to(device).long()375 376 _, batch_speaker_embeddings, _, _, _ = ssl_model.forward_for_export(377 input_signal=audio_segmented.to(device), input_signal_length=audio_seg_len, normalize_content=True,378 )379 380 for idx in range(batch['audio'].shape[0]):381 _speaker = batch['speaker'][idx].item()382 wav_path = batch['wav_path'][idx]383 384 wav_id = batch['rel_audio_path_as_text_id'][idx]385 wav_and_id_list.append((wav_path, wav_id, _speaker))386 content_embedding = batch_content_embedding[idx].detach()387 content_log_probs = batch_content_log_probs[:, idx, :].detach() # (content lob prob is (t, b, c))388 encoded_len = batch_encoded_len[idx].detach()389 content_embedding = content_embedding[: encoded_len.item()]390 content_embedding = content_embedding.t()391 content_log_probs = content_log_probs[: encoded_len.item()]392 content_log_probs = content_log_probs.t()393 content_probs = torch.exp(content_log_probs)394 395 duration = torch.ones(content_embedding.shape[1]) * args.ssl_downsampling_factor396 397 bsi_start = segment_indices[idx][0]398 bsi_end = segment_indices[idx][1]399 speaker_embedding = torch.mean(batch_speaker_embeddings[bsi_start : bsi_end + 1], dim=0)400 401 l2_norm = torch.norm(speaker_embedding, p=2)402 speaker_embedding = speaker_embedding / l2_norm403 404 if args.ssl_content_emb_type == "probs":405 # content embedding is only character probabilities406 final_content_embedding = content_probs407 elif args.ssl_content_emb_type == "embedding":408 # content embedding is only output of content head of SSL backbone409 final_content_embedding = content_embedding410 elif args.ssl_content_emb_type == "log_probs":411 # content embedding is only log of character probabilities412 final_content_embedding = content_log_probs413 elif args.ssl_content_emb_type == "embedding_and_probs":414 # content embedding is the concatenation of character probabilities and output of content head of SSL backbone415 final_content_embedding = torch.cat([content_embedding, content_probs], dim=0)416 417 if args.use_unique_tokens == 1:418 # group content embeddings with same predicted token (by averaging) and add the durations of the grouped embeddings419 # Eg. By default each content embedding corresponds to 4 frames of spectrogram (ssl_downsampling_factor)420 # If we group 3 content embeddings, the duration of the grouped embedding will be 12 frames.421 # This is useful for adapting the duration during inference based on the speaker.422 token_predictions = torch.argmax(content_probs, dim=0)423 content_buffer = [final_content_embedding[:, 0]]424 unique_content_embeddings = []425 unique_tokens = []426 durations = []427 for _t in range(1, final_content_embedding.shape[1]):428 if token_predictions[_t] == token_predictions[_t - 1]:429 content_buffer.append(final_content_embedding[:, _t])430 else:431 durations.append(len(content_buffer) * args.ssl_downsampling_factor)432 unique_content_embeddings.append(torch.mean(torch.stack(content_buffer), dim=0))433 content_buffer = [final_content_embedding[:, _t]]434 unique_tokens.append(token_predictions[_t].item())435 436 if len(content_buffer) > 0:437 durations.append(len(content_buffer) * args.ssl_downsampling_factor)438 unique_content_embeddings.append(torch.mean(torch.stack(content_buffer), dim=0))439 unique_tokens.append(token_predictions[_t].item())440 441 unique_content_embedding = torch.stack(unique_content_embeddings)442 final_content_embedding = unique_content_embedding.t()443 duration = torch.tensor(durations).float()444 445 mel_len = int(batch['audio_len'][idx].item() / stft_params['hop_length'])446 item_mel = batch_mel_specs[idx][:, :mel_len]447 448 wav_text_id = batch["rel_audio_path_as_text_id"][idx]449 content_emb_fn = f"{args.ssl_content_emb_type}_content_embedding_{wav_text_id}.pt"450 speaker_emb_fn = f"speaker_embedding_{wav_text_id}.pt"451 duration_fn = f"duration_embedding_{wav_text_id}.pt" # embedding just for namesake452 content_emb_fp = os.path.join(dataset.sup_data_dir, content_emb_fn)453 speaker_emb_fp = os.path.join(dataset.sup_data_dir, speaker_emb_fn)454 duration_fp = os.path.join(dataset.sup_data_dir, duration_fn)455 456 mel_spec_fn = f"mel_spec_{wav_text_id}.pt"457 mel_spec_fp = os.path.join(dataset.sup_data_dir, mel_spec_fn)458 459 torch.save(item_mel.cpu(), mel_spec_fp)460 torch.save(final_content_embedding.cpu(), content_emb_fp)461 torch.save(speaker_embedding.cpu(), speaker_emb_fp)462 torch.save(duration.cpu(), duration_fp)463 464 et = time.time()465 logging.info(466 "Processed Batch {} of {} | Time per batch: {:.4f} s".format(467 bidx + 1, len(dataloader), (et - st) / bidx468 )469 )470 471 if args.compute_pitch_contours == 1:472 speaker_wise_records = {}473 for row in wav_and_id_list:474 wav_path, wav_id, speaker = row475 if speaker not in speaker_wise_records:476 speaker_wise_records[speaker] = []477 speaker_wise_records[speaker].append(478 {479 "wav_path": wav_path,480 "wav_id": wav_id,481 "sup_data_dir": dataset.sup_data_dir,482 "stft_params": stft_params,483 "speaker": speaker,484 }485 )486 487 filtered_records = []488 for speaker in speaker_wise_records:489 if args.num_pitch_per_speaker is not None:490 filtered_records += speaker_wise_records[speaker][: args.num_pitch_per_speaker]491 else:492 filtered_records += speaker_wise_records[speaker]493 494 with Pool(args.pool_workers) as p:495 p.map(save_pitch_contour, filtered_records)496 497 compute_pitch_stats(filtered_records)498 499 500if __name__ == '__main__':501 main()502 