import multiprocessing from threading import Thread from queue import Queue import sys from contextlib import contextmanager class _BeatProcessorWorker: def __init__(self): self.q_new = Queue() self.q_ready = Queue() self.worker_thd = Thread(target=self._worker_thread_entry, daemon=True) self.worker_thd.start() self.q_new.put(1) def _worker_thread_entry(self): from madmom.features import DBNDownBeatTrackingProcessor, RNNDownBeatProcessor while True: item = self.q_new.get() if item is None: break try: beat_processor = RNNDownBeatProcessor(fps=100) down_beat_processor = DBNDownBeatTrackingProcessor( beats_per_bar=[3, 4], fps=100 ) self.q_ready.put((beat_processor, down_beat_processor)) except Exception as e: print(f"error constructing madmom resources: {e}", file=sys.stderr) sys.exit(1) @contextmanager def processors(self): p = self.q_ready.get() try: yield p finally: self.q_new.put(1) _global_beat_processor_worker = None def _create_worker(): global _global_beat_processor_worker if _global_beat_processor_worker is None: _global_beat_processor_worker = _BeatProcessorWorker() def _detect_beats(wav_path): with _global_beat_processor_worker.processors() as ( beat_processor, down_beat_processor, ): dist = beat_processor.process(str(wav_path)) return down_beat_processor.process(dist) class BeatProcessor: def __init__(self, num_workers=1): self.pool = multiprocessing.Pool(num_workers, initializer=_create_worker) def close(self): self.pool.close() self.pool.join() def detect_beats(self, wav_path): return self.pool.apply(_detect_beats, (str(wav_path),))