124 lines
4.8 KiB
Python
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()
|