seprq 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- seprq/__init__.py +86 -0
- seprq-0.1.0.dist-info/METADATA +45 -0
- seprq-0.1.0.dist-info/RECORD +5 -0
- seprq-0.1.0.dist-info/WHEEL +5 -0
- seprq-0.1.0.dist-info/top_level.txt +1 -0
seprq/__init__.py
ADDED
|
@@ -0,0 +1,86 @@
|
|
|
1
|
+
"""seprq -- SepRQ / BEST-RQ (50 Hz) SSL speech feature extractors.
|
|
2
|
+
|
|
3
|
+
The CNN frontend + 12-layer Conformer encoder and the global norm stats are
|
|
4
|
+
downloaded from the Hugging Face Hub (SevKod/SepRQ) on first use; only the
|
|
5
|
+
requested model is fetched. Uses upstream speechbrain (>=1.0.3) only.
|
|
6
|
+
|
|
7
|
+
from seprq import SepRQEncoder
|
|
8
|
+
speech_encoder = SepRQEncoder("SepRQ") # or "BestRQ_50Hz"
|
|
9
|
+
feats = speech_encoder("audio.wav") # path -> [1, T, 576]
|
|
10
|
+
"""
|
|
11
|
+
import torch
|
|
12
|
+
from hyperpyyaml import load_hyperpyyaml
|
|
13
|
+
from huggingface_hub import hf_hub_download
|
|
14
|
+
|
|
15
|
+
SAMPLE_RATE = 16000
|
|
16
|
+
|
|
17
|
+
__all__ = ["SepRQEncoder", "MODELS", "REPO_ID"]
|
|
18
|
+
__version__ = "0.1.0"
|
|
19
|
+
|
|
20
|
+
REPO_ID = "SevKod/SepRQ"
|
|
21
|
+
|
|
22
|
+
# model name -> (subfolder in the repo, normalize-checkpoint filename)
|
|
23
|
+
MODELS = {
|
|
24
|
+
"SepRQ": ("SepRQ/reduced_scale/2_streams", "normalize.ckpt"),
|
|
25
|
+
"BestRQ_50Hz": ("BestRQ_50Hz/reduced_scale", "normalize_running_stats.ckpt"),
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class SepRQEncoder(torch.nn.Module):
|
|
30
|
+
def __init__(self, model="SepRQ", repo_id=REPO_ID, device=None):
|
|
31
|
+
super().__init__()
|
|
32
|
+
if model not in MODELS:
|
|
33
|
+
raise ValueError(f"model must be one of {list(MODELS)}, got {model!r}")
|
|
34
|
+
subdir, norm_name = MODELS[model]
|
|
35
|
+
self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
|
|
36
|
+
|
|
37
|
+
# Download ONLY the files needed for this model.
|
|
38
|
+
yaml_path = hf_hub_download(repo_id, "seprq_inference.yaml")
|
|
39
|
+
model_ckpt = hf_hub_download(repo_id, f"{subdir}/model.ckpt")
|
|
40
|
+
norm_ckpt = hf_hub_download(repo_id, f"{subdir}/{norm_name}")
|
|
41
|
+
|
|
42
|
+
with open(yaml_path) as f:
|
|
43
|
+
hp = load_hyperpyyaml(f)
|
|
44
|
+
|
|
45
|
+
self.melspec = hp["compute_features"].to(self.device)
|
|
46
|
+
self.cnn = hp["CNN"].to(self.device)
|
|
47
|
+
self.wrapper = hp["wrapper"].to(self.device) # 12-layer Conformer
|
|
48
|
+
|
|
49
|
+
# model.ckpt is the state_dict of ModuleList([CNN, wrapper]) -> 0.*/1.*
|
|
50
|
+
torch.nn.ModuleList([self.cnn, self.wrapper]).load_state_dict(
|
|
51
|
+
torch.load(model_ckpt, map_location=self.device)
|
|
52
|
+
)
|
|
53
|
+
# Global norm stats: the checkpoint stores running_mean / running_var.
|
|
54
|
+
norm = torch.load(norm_ckpt, map_location=self.device)
|
|
55
|
+
self.running_mean = norm["running_mean"].to(self.device)
|
|
56
|
+
self.running_std = torch.sqrt(norm["running_var"].to(self.device) + 1e-5)
|
|
57
|
+
|
|
58
|
+
self.cnn.eval()
|
|
59
|
+
self.wrapper.eval()
|
|
60
|
+
|
|
61
|
+
@staticmethod
|
|
62
|
+
def _load(path):
|
|
63
|
+
"""Load an audio file -> mono 16 kHz float32 waveform [samples]."""
|
|
64
|
+
import soundfile as sf
|
|
65
|
+
|
|
66
|
+
data, sr = sf.read(path, dtype="float32", always_2d=True) # [samp, ch]
|
|
67
|
+
wav = torch.from_numpy(data).mean(1) # mono
|
|
68
|
+
if sr != SAMPLE_RATE:
|
|
69
|
+
import torchaudio
|
|
70
|
+
|
|
71
|
+
wav = torchaudio.functional.resample(wav, sr, SAMPLE_RATE)
|
|
72
|
+
return wav
|
|
73
|
+
|
|
74
|
+
@torch.no_grad()
|
|
75
|
+
def forward(self, audio):
|
|
76
|
+
"""audio: path to an audio file, or a 1D 16 kHz waveform (array/tensor).
|
|
77
|
+
Returns SSL features [1, T, 576] from the final Conformer layer.
|
|
78
|
+
Call the instance directly: ``speech_encoder(audio)``."""
|
|
79
|
+
if isinstance(audio, (str, bytes)) or hasattr(audio, "__fspath__"):
|
|
80
|
+
audio = self._load(audio)
|
|
81
|
+
wavs = torch.as_tensor(audio, dtype=torch.float32).reshape(1, -1)
|
|
82
|
+
wavs = wavs.to(self.device)
|
|
83
|
+
wav_lens = torch.tensor([1.0], device=self.device)
|
|
84
|
+
feats = self.melspec(wavs) # Fbank [1, T, 80]
|
|
85
|
+
feats = (feats - self.running_mean) / self.running_std # global norm
|
|
86
|
+
return self.wrapper(self.cnn(feats), wav_lens) # [1, T, 576]
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: seprq
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: SepRQ / BEST-RQ (50 Hz) SSL speech feature extractors (weights hosted on the Hugging Face Hub).
|
|
5
|
+
Author: SevKod
|
|
6
|
+
License: Apache-2.0
|
|
7
|
+
Project-URL: Homepage, https://huggingface.co/SevKod/SepRQ
|
|
8
|
+
Requires-Python: >=3.9
|
|
9
|
+
Description-Content-Type: text/markdown
|
|
10
|
+
Requires-Dist: speechbrain>=1.0.3
|
|
11
|
+
Requires-Dist: hyperpyyaml
|
|
12
|
+
Requires-Dist: huggingface_hub
|
|
13
|
+
Requires-Dist: torch
|
|
14
|
+
Requires-Dist: soundfile
|
|
15
|
+
|
|
16
|
+
# seprq
|
|
17
|
+
|
|
18
|
+
SSL speech feature extractors for **SepRQ** and **BEST-RQ (50 Hz)**. The
|
|
19
|
+
encoder (CNN frontend + 12-layer Conformer) and global norm stats are pulled
|
|
20
|
+
from the Hugging Face Hub ([SevKod/SepRQ](https://huggingface.co/SevKod/SepRQ))
|
|
21
|
+
on first use — only the requested model is downloaded. Upstream
|
|
22
|
+
`speechbrain>=1.0.3` only; runs on CPU or GPU.
|
|
23
|
+
|
|
24
|
+
## Install
|
|
25
|
+
|
|
26
|
+
```bash
|
|
27
|
+
pip install git+https://github.com/SevKod/SepRQ.git # install straight from GitHub
|
|
28
|
+
# or, from a local checkout:
|
|
29
|
+
pip install .
|
|
30
|
+
```
|
|
31
|
+
|
|
32
|
+
## Use
|
|
33
|
+
|
|
34
|
+
```python
|
|
35
|
+
from seprq import SepRQEncoder
|
|
36
|
+
|
|
37
|
+
speech_encoder = SepRQEncoder("SepRQ") # or "BestRQ_50Hz"
|
|
38
|
+
feats = speech_encoder("utterance.wav") # torch.Tensor [1, T, 576]
|
|
39
|
+
```
|
|
40
|
+
|
|
41
|
+
Calling the encoder runs `forward`. It accepts a **file path** (any
|
|
42
|
+
format/sample rate — decoded, downmixed to mono and resampled to 16 kHz for
|
|
43
|
+
you) or a raw 1D 16 kHz waveform (numpy array / tensor).
|
|
44
|
+
|
|
45
|
+
The repo is private, so authenticate once: `huggingface-cli login` (or set `HF_TOKEN`).
|
|
@@ -0,0 +1,5 @@
|
|
|
1
|
+
seprq/__init__.py,sha256=0cL2A9B9wN6FQl9Na-XG7CfZkKCxPSjze-JndpAeRGg,3616
|
|
2
|
+
seprq-0.1.0.dist-info/METADATA,sha256=h1oSpvoLTOPswG7fyacdBT3_2VQEOOeOuim-wJICMY8,1444
|
|
3
|
+
seprq-0.1.0.dist-info/WHEEL,sha256=YVMoNqKzERt-wjUZwJ33xBGAwnFl-4cqbYkTtWa4itE,91
|
|
4
|
+
seprq-0.1.0.dist-info/top_level.txt,sha256=aj5lE6fu1h5y7OvEl6bOJSRdk2y0JlyHUXJsBpil_Ao,6
|
|
5
|
+
seprq-0.1.0.dist-info/RECORD,,
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
seprq
|