Views
No views yet
1import pathlib as pl
2import numpy as np
3import matplotlib.pyplot as plt
4import torch
5import torchaudio
6from speechbrain.inference.vocoders import UnitHIFIGAN
7from speechbrain.lobes.models.huggingface_transformers import (
8 hubert,
9 wav2vec2,
10 wavlm,
11)
12from speechbrain.lobes.models.huggingface_transformers.discrete_ssl import (
13 DiscreteSSL,
14)
15
16ENCODER_CLASSES = {
17 "HuBERT": hubert.HuBERT,
18 "Wav2Vec2": wav2vec2.Wav2Vec2,
19 "WavLM": wavlm.WavLM,
20}
21
22kmeans_folder = "poonehmousavi/SSL_Quantization"
23kmeans_dataset = "LJSpeech" # LibriSpeech-100-360-500
24num_clusters = 1000
25encoder_type = "HuBERT" # one of [HuBERT, Wav2Vec2, WavLM]
26encoder_source = "facebook/hubert-large-ll60k"
27layer = [3, 7, 12, 18, 23]
28vocoder_source = (
29 "chaanks/hifigan-unit-hubert-ll60k-l3-7-12-18-23-k1000-ljspeech-ljspeech"
30)
31save_path = pl.Path(".tmpdir")
32device = "cuda"
33sample_rate = 16000
34
35wav = "chaanks/hifigan-unit-hubert-ll60k-l3-7-12-18-23-k1000-ljspeech-ljspeech/test.wav"
36
37encoder_class = ENCODER_CLASSES[encoder_type]
38encoder = encoder_class(
39 source=encoder_source,
40 save_path=(save_path / "encoder").as_posix(),
41 output_norm=False,
42 freeze=True,
43 freeze_feature_extractor=True,
44 apply_spec_augment=False,
45 output_all_hiddens=True,
46).to(device)
47
48discrete_encoder = DiscreteSSL(
49 save_path=(save_path / "discrete_encoder").as_posix(),
50 ssl_model=encoder,
51 kmeans_dataset=kmeans_dataset,
52 kmeans_repo_id=kmeans_folder,
53 num_clusters=num_clusters,
54)
55
56vocoder = UnitHIFIGAN.from_hparams(
57 source=vocoder_source,
58 run_opts={"device": str(device)},
59 savedir=(save_path / "vocoder").as_posix(),
60)
61
62audio = vocoder.load_audio(wav)
63audio = audio.unsqueeze(0).to(device)
64
65deduplicates = [False for _ in layer]
66bpe_tokenizers = [None for _ in layer]
67tokens, _, _ = discrete_encoder(
68 audio,
69 SSL_layers=layer,
70 deduplicates=deduplicates,
71 bpe_tokenizers=bpe_tokenizers,
72)
73tokens = tokens.cpu().squeeze(0)
74
75num_layer = len(layer)
76offsets = torch.arange(num_layer) * num_clusters
77tokens = tokens + offsets
78
79waveform = vocoder.decode_unit(tokens)
80torchaudio.save("pred.wav", waveform.cpu(), sample_rate=sample_rate)