Add lazy piano transcription and optional leased server worker
This commit is contained in:
123
scripts/piano/worker.py
Normal file
123
scripts/piano/worker.py
Normal file
@@ -0,0 +1,123 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user