CoolFace
Modelpublic

RASMUS/Finnish-ASR-Canary-v2

sourceHugging Facemitupdated 7mo agoView on Hugging Face
0likes1.2kdownloads
make_supdata.py502 linesDownload Raw Back to ssl_tts
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