Files
ytplayer/scripts/piano/worker.py

124 lines
4.8 KiB
Python

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