50 lines
1.6 KiB
Python
50 lines
1.6 KiB
Python
"""One-time download of the speechbrain enhancement model.
|
|
|
|
Each worker loads the model from the local cache, so run this once before
|
|
starting all 20 workers:
|
|
|
|
.venv/bin/python prefetch_model.py
|
|
"""
|
|
|
|
import logging
|
|
import sys
|
|
|
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
|
log = logging.getLogger("prefetch")
|
|
|
|
|
|
def main() -> int:
|
|
from config import load_config
|
|
cfg = load_config()
|
|
import torch
|
|
from speechbrain.inference.enhancement import SpectralMaskEnhancement
|
|
|
|
log.info("loading %s (first run downloads weights) ...", cfg.enhance_model)
|
|
model = SpectralMaskEnhancement.from_hparams(
|
|
source=cfg.enhance_model,
|
|
savedir="/root/.cache/speechbrain-enhancement",
|
|
)
|
|
model.mods.enhance_model.eval()
|
|
|
|
# quick self-test: 0.5 s of 16 kHz sine + noise
|
|
import numpy as np
|
|
n = int(cfg.audio_rate * 0.5)
|
|
t = np.arange(n) / cfg.audio_rate
|
|
x = (0.3 * np.sin(2 * np.pi * 440 * t) + 0.05 * np.random.randn(n))
|
|
x = x / np.abs(x).max()
|
|
with torch.no_grad():
|
|
w = torch.from_numpy(x.astype(np.float32)).unsqueeze(0).to(model.device)
|
|
# lengths are relative (1.0 = full sequence); older speechbrain
|
|
# builds don't accept the kwarg at all -> TypeError fallback.
|
|
try:
|
|
out = model.enhance_batch(w, lengths=torch.tensor([1.0]))
|
|
except TypeError:
|
|
out = model.enhance_batch(w)
|
|
log.info("self-test OK: in %d samples -> out %d samples",
|
|
n, out.numel())
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
sys.exit(main())
|