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 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,5 @@
1
+ Wheel-Version: 1.0
2
+ Generator: setuptools (84.0.0)
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
5
+
@@ -0,0 +1 @@
1
+ seprq