Files
livekit-cameras/prefetch_model.py

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())