fosters commited on
Commit
a845faf
·
verified ·
1 Parent(s): 3d9eb32

Upload app.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. app.py +9 -14
app.py CHANGED
@@ -8,40 +8,36 @@ import numpy as np
8
  import pandas as pd
9
  import requests
10
  import soundfile as sf
11
- import torch
12
  import torchaudio
 
13
  from sklearn.cluster import AgglomerativeClustering
14
- from transformers import Wav2Vec2FeatureExtractor, WavLMForXVector
 
15
 
16
  os.environ.setdefault("HF_XET_HIGH_PERFORMANCE", "1")
17
 
18
- # 1 PyTorch thread per worker — lets N_CPUS threads run inference in parallel
19
- # instead of one wide inference that blocks everyone else.
20
- torch.set_num_threads(1)
21
-
22
  N_CPUS = os.cpu_count() or 2
23
 
24
  DATASETS_SERVER = "https://datasets-server.huggingface.co"
25
- MODEL_ID = "microsoft/wavlm-base-plus-sv"
26
  TARGET_SR = 16000
27
 
28
  _feature_extractor = None
29
  _model = None
30
- _init_lock = threading.Lock() # only for one-time model init
31
 
32
 
33
  def _load_model():
34
  global _feature_extractor, _model
35
  with _init_lock:
36
  if _model is None:
37
- _feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(MODEL_ID)
38
- _model = WavLMForXVector.from_pretrained(MODEL_ID)
39
- _model.eval()
40
  return _feature_extractor, _model
41
 
42
 
43
  def _embed(audio_array: np.ndarray, sr: int, max_sec: int) -> np.ndarray:
44
- # Thread-safe: eval() + no_grad() — weights are read-only, GIL released in C++ ops
45
  fe, mdl = _load_model()
46
  waveform = torch.tensor(audio_array, dtype=torch.float32)
47
  if waveform.ndim == 2:
@@ -50,8 +46,7 @@ def _embed(audio_array: np.ndarray, sr: int, max_sec: int) -> np.ndarray:
50
  waveform = torchaudio.functional.resample(waveform, sr, TARGET_SR)
51
  waveform = waveform[: max_sec * TARGET_SR]
52
  inputs = fe(waveform.numpy(), sampling_rate=TARGET_SR, return_tensors="pt")
53
- with torch.no_grad():
54
- out = mdl(**inputs)
55
  return out.embeddings.squeeze().numpy()
56
 
57
 
 
8
  import pandas as pd
9
  import requests
10
  import soundfile as sf
 
11
  import torchaudio
12
+ import torch
13
  from sklearn.cluster import AgglomerativeClustering
14
+ from optimum.onnxruntime import ORTModelForAudioXVector
15
+ from transformers import Wav2Vec2FeatureExtractor
16
 
17
  os.environ.setdefault("HF_XET_HIGH_PERFORMANCE", "1")
18
 
 
 
 
 
19
  N_CPUS = os.cpu_count() or 2
20
 
21
  DATASETS_SERVER = "https://datasets-server.huggingface.co"
22
+ ONNX_MODEL_ID = "fosters/wavlm-base-plus-sv-onnx"
23
  TARGET_SR = 16000
24
 
25
  _feature_extractor = None
26
  _model = None
27
+ _init_lock = threading.Lock()
28
 
29
 
30
  def _load_model():
31
  global _feature_extractor, _model
32
  with _init_lock:
33
  if _model is None:
34
+ _feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(ONNX_MODEL_ID)
35
+ # ONNX Runtime: thread-safe, no GIL concerns
36
+ _model = ORTModelForAudioXVector.from_pretrained(ONNX_MODEL_ID)
37
  return _feature_extractor, _model
38
 
39
 
40
  def _embed(audio_array: np.ndarray, sr: int, max_sec: int) -> np.ndarray:
 
41
  fe, mdl = _load_model()
42
  waveform = torch.tensor(audio_array, dtype=torch.float32)
43
  if waveform.ndim == 2:
 
46
  waveform = torchaudio.functional.resample(waveform, sr, TARGET_SR)
47
  waveform = waveform[: max_sec * TARGET_SR]
48
  inputs = fe(waveform.numpy(), sampling_rate=TARGET_SR, return_tensors="pt")
49
+ out = mdl(**inputs)
 
50
  return out.embeddings.squeeze().numpy()
51
 
52