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

Upload app.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. app.py +13 -7
app.py CHANGED
@@ -15,18 +15,24 @@ from transformers import Wav2Vec2FeatureExtractor, WavLMForXVector
15
 
16
  os.environ.setdefault("HF_XET_HIGH_PERFORMANCE", "1")
17
 
 
 
 
 
 
 
18
  DATASETS_SERVER = "https://datasets-server.huggingface.co"
19
  MODEL_ID = "microsoft/wavlm-base-plus-sv"
20
  TARGET_SR = 16000
21
 
22
  _feature_extractor = None
23
  _model = None
24
- _embed_lock = threading.Lock() # WavLM inference is not thread-safe
25
 
26
 
27
  def _load_model():
28
  global _feature_extractor, _model
29
- with _embed_lock:
30
  if _model is None:
31
  _feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(MODEL_ID)
32
  _model = WavLMForXVector.from_pretrained(MODEL_ID)
@@ -35,6 +41,7 @@ def _load_model():
35
 
36
 
37
  def _embed(audio_array: np.ndarray, sr: int, max_sec: int) -> np.ndarray:
 
38
  fe, mdl = _load_model()
39
  waveform = torch.tensor(audio_array, dtype=torch.float32)
40
  if waveform.ndim == 2:
@@ -43,9 +50,8 @@ def _embed(audio_array: np.ndarray, sr: int, max_sec: int) -> np.ndarray:
43
  waveform = torchaudio.functional.resample(waveform, sr, TARGET_SR)
44
  waveform = waveform[: max_sec * TARGET_SR]
45
  inputs = fe(waveform.numpy(), sampling_rate=TARGET_SR, return_tensors="pt")
46
- with _embed_lock:
47
- with torch.no_grad():
48
- out = mdl(**inputs)
49
  return out.embeddings.squeeze().numpy()
50
 
51
 
@@ -186,8 +192,8 @@ def identify_speakers(
186
  errors: list[str] = []
187
  done = 0
188
 
189
- # Process all repos in parallel (downloads are I/O bound)
190
- with concurrent.futures.ThreadPoolExecutor(max_workers=8) as ex:
191
  future_to_repo = {
192
  ex.submit(_process_repo, repo, int(samples_per_book), int(audio_sec), token): repo
193
  for repo in repos
 
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)
 
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
  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
 
 
192
  errors: list[str] = []
193
  done = 0
194
 
195
+ # Process repos in parallel capped at N_CPUS since embed is the bottleneck
196
+ with concurrent.futures.ThreadPoolExecutor(max_workers=N_CPUS) as ex:
197
  future_to_repo = {
198
  ex.submit(_process_repo, repo, int(samples_per_book), int(audio_sec), token): repo
199
  for repo in repos