"""Optional CPU piano worker. Polls authenticated jobs; loads model on first job.""" import json import os import re import tempfile import threading import time import urllib.request import urllib.error from pathlib import Path BASE = os.environ.get('YTP_BASE', 'http://ytplayer:3000').rstrip('/') TOKEN = os.environ.get('YTP_TOKEN', '') CHECKPOINT = Path(os.environ.get('PIANO_CHECKPOINT', '/models/piano.pth')) MODEL_URL = 'https://zenodo.org/record/4034264/files/CRNN_note_F1%3D0.9677_pedal_F1%3D0.9186.pth?download=1' ID = re.compile(r'^(?:[\w-]{11}|upl_[a-f0-9]{12})$') def api(path, body): request = urllib.request.Request(BASE + path, data=json.dumps(body).encode(), headers={'Authorization': 'Bearer ' + TOKEN, 'Content-Type': 'application/json'}) with urllib.request.urlopen(request, timeout=45) as response: return json.load(response) def compact(events): return [{'pitch': int(event['midi_note']), 'start': round(float(event['onset_time']), 3), 'end': round(float(event['offset_time']), 3), 'velocity': max(1, min(127, int(event['velocity']))), 'hand': 'left' if int(event['midi_note']) < 60 else 'right'} for event in events if event['offset_time'] > event['onset_time']] def download(url, destination, limit): with urllib.request.urlopen(url, timeout=60) as response, open(destination, 'wb') as target: size = 0 while True: chunk = response.read(1024 * 1024) if not chunk: break size += len(chunk) if size > limit: raise ValueError('Download exceeds worker limit') target.write(chunk) def work(job, model): video = job['videoId'] if not ID.fullmatch(video): raise ValueError('Invalid job id') path = '/api/uploads/' + video if video.startswith('upl_') else '/api/media/' + video + '?a=1' if job['audioPath'] != path: raise ValueError('Invalid audio path') endpoint = '/api/piano-worker/' + video stop = threading.Event() def heartbeat(): while not stop.wait(30): try: api(endpoint, {'action': 'heartbeat', 'lease': job['lease']}) except Exception: # Completion still checks the lease; an expired worker cannot publish. pass thread = threading.Thread(target=heartbeat, daemon=True) thread.start() try: if model is None: if not CHECKPOINT.exists() or CHECKPOINT.stat().st_size < 160000000: CHECKPOINT.parent.mkdir(parents=True, exist_ok=True) temporary = CHECKPOINT.with_suffix('.part') download(MODEL_URL, temporary, 250000000) if temporary.stat().st_size < 160000000: raise ValueError('Incomplete model checkpoint') temporary.replace(CHECKPOINT) import torch from piano_transcription_inference import PianoTranscription torch.set_num_threads(int(os.environ.get('PIANO_THREADS', '2'))) model = PianoTranscription(device='cpu', checkpoint_path=str(CHECKPOINT)) from piano_transcription_inference import load_audio, sample_rate with tempfile.TemporaryDirectory() as folder: media = Path(folder) / 'audio.bin' try: download(BASE + path, media, 100 * 1048576) except urllib.error.HTTPError as error: if error.code != 404 or video.startswith('upl_'): raise download(BASE + path.split('?')[0], media, 100 * 1048576) audio, _ = load_audio(str(media), sr=sample_rate, mono=True) if len(audio) / sample_rate > 900: raise ValueError('Worker audio is limited to 15 minutes') result = model.transcribe(audio, None) notes = compact(result['est_note_events']) if len(notes) > 30000: raise ValueError('Too many notes') api(endpoint, {'action': 'complete', 'lease': job['lease'], 'notes': notes}) return model except Exception as error: api(endpoint, {'action': 'fail', 'lease': job['lease'], 'error': str(error)[:500]}) return model finally: stop.set() thread.join(timeout=1) def main(): if len(TOKEN) < 24: print('Piano worker disabled: configure a dedicated worker token.', flush=True) while True: time.sleep(300) model = None while True: try: job = api('/api/piano-worker/claim', {}).get('job') if job: model = work(job, model) else: time.sleep(10) except Exception as error: print('Piano worker: ' + str(error), flush=True) time.sleep(15) if __name__ == '__main__': main()