Add lazy piano transcription and optional leased server worker

This commit is contained in:
Jonathan Sykes
2026-10-03 15:47:53 +08:00
parent d4354d220e
commit 0374e19396
20 changed files with 1107 additions and 1 deletions

9
scripts/piano/Dockerfile Normal file
View File

@@ -0,0 +1,9 @@
FROM python:3.10-slim
ENV PYTHONUNBUFFERED=1 PIP_NO_CACHE_DIR=1 OMP_NUM_THREADS=2
RUN apt-get update && apt-get install -y --no-install-recommends ffmpeg libsndfile1 wget && rm -rf /var/lib/apt/lists/*
WORKDIR /app
COPY requirements.txt .
RUN pip install -r requirements.txt
COPY worker.py .
VOLUME ["/models"]
CMD ["python", "worker.py"]

View File

@@ -0,0 +1,9 @@
--extra-index-url https://download.pytorch.org/whl/cpu
torch==2.5.1+cpu
piano-transcription-inference==0.0.6
numpy==1.23.5
librosa==0.9.2
numba==0.57.1
scipy==1.10.1
soundfile==0.12.1
mido==1.3.3

View File

@@ -0,0 +1,18 @@
import unittest
from worker import compact
class CompactNotesTest(unittest.TestCase):
def test_hand_split_and_velocity(self):
notes = compact([{'midi_note': 48, 'onset_time': 1.23456, 'offset_time': 1.78, 'velocity': 85}, {'midi_note': 60, 'onset_time': 2, 'offset_time': 3, 'velocity': 200}])
self.assertEqual(notes[0], {'pitch': 48, 'start': 1.235, 'end': 1.78, 'velocity': 85, 'hand': 'left'})
self.assertEqual(notes[1]['hand'], 'right')
self.assertEqual(notes[1]['velocity'], 127)
def test_empty_or_zero_duration(self):
self.assertEqual(compact([]), [])
self.assertEqual(compact([{'midi_note': 60, 'onset_time': 2, 'offset_time': 2, 'velocity': 85}]), [])
if __name__ == '__main__':
unittest.main()

123
scripts/piano/worker.py Normal file
View 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()