Keep TTS and STT workers resident between requests
This commit is contained in:
@@ -0,0 +1,59 @@
|
||||
"""Synthetic process tests: audio workers survive a second request and stop cleanly."""
|
||||
import json
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
from inference import Scheduler
|
||||
from stt import STT
|
||||
from test_stt import audio
|
||||
from tts_test import TTSTests
|
||||
|
||||
def done(manager):
|
||||
for _ in range(200):
|
||||
job=manager.status()['job']
|
||||
if job and job['state']!='running':return job
|
||||
time.sleep(.05)
|
||||
raise AssertionError('synthetic worker timed out')
|
||||
|
||||
class ResidencyTests(unittest.TestCase):
|
||||
def test_tts_reuses_process_and_gpu_switch_eviction(self):
|
||||
with tempfile.TemporaryDirectory() as d:
|
||||
script=Path(d)/'worker.py'
|
||||
script.write_text('import json,sys,pathlib\nprint(json.dumps({"ready":True}),flush=True)\nfor line in sys.stdin:\n q=json.loads(line);p=pathlib.Path(q["output"]);(p/"result.wav").write_bytes(b"synthetic wav");(p/"result.json").write_text(json.dumps({"ok":True}))\n')
|
||||
runtime=SimpleNamespace(paths=lambda:(Path(sys.executable),Path(d)/'model'),status=lambda:{'installed':True})
|
||||
profile=dict(id='p',kind='audio',runnable=True,parameters={'speed':1})
|
||||
manager=TTSTests(Path(d)/'jobs',SimpleNamespace(status=lambda:{'profiles':[profile]}),runtime)
|
||||
worker=SimpleNamespace(stop=lambda:None)
|
||||
scheduler=Scheduler(worker);scheduler.evict_tts=manager.unload_idle
|
||||
manager.acquire=scheduler.tts_reservation
|
||||
def gpu():return [dict(uuid='synthetic-3060',name='RTX 3060',processes=int(bool(manager.process and manager.process.poll() is None)),free_mib=11000)]
|
||||
try:
|
||||
with patch('tts_test.WORKER',script),patch('tts_test.probe',side_effect=gpu),patch('tts_test.cgroup_headroom',return_value=None):
|
||||
with scheduler.lease(('chat','medium')):pass
|
||||
manager.start('p','first');self.assertEqual(done(manager)['state'],'complete');pid=manager.process.pid
|
||||
manager.start('p','second');self.assertEqual(done(manager)['state'],'complete');self.assertEqual(manager.process.pid,pid)
|
||||
with scheduler.lease(('chat','medium')):self.assertTrue(manager.status()['loaded'])
|
||||
with scheduler.lease(('image',)):self.assertFalse(manager.status()['loaded'])
|
||||
finally:manager.stop()
|
||||
|
||||
def test_stt_reuses_cpu_server_until_stop(self):
|
||||
with tempfile.TemporaryDirectory() as d:
|
||||
script=Path(d)/'server.py'
|
||||
script.write_text('#!'+sys.executable+'\nimport argparse,json\nfrom http.server import BaseHTTPRequestHandler,HTTPServer\np=argparse.ArgumentParser();p.add_argument("--port",type=int);a,_=p.parse_known_args()\nclass H(BaseHTTPRequestHandler):\n def log_message(self,*x):pass\n def do_GET(self):\n self.send_response(200);self.end_headers();self.wfile.write(b"ok")\n def do_POST(self):\n self.rfile.read(int(self.headers["Content-Length"]));b=json.dumps({"text":"synthetic transcript"}).encode();self.send_response(200);self.send_header("Content-Length",str(len(b)));self.end_headers();self.wfile.write(b)\nHTTPServer(("127.0.0.1",a.port),H).serve_forever()\n')
|
||||
script.chmod(0o700)
|
||||
profile=dict(id='p',kind='stt',runnable=True,model_id='m')
|
||||
catalog=SimpleNamespace(root=Path(d),status=lambda:{'entries':[]})
|
||||
manager=STT(SimpleNamespace(status=lambda:{'profiles':[profile]}),catalog,SimpleNamespace(status=lambda:{'active':None}))
|
||||
try:
|
||||
with patch.object(manager,'build',return_value=script),patch.object(manager,'projector',return_value={'id':'projector'}),patch('stt.cgroup_headroom',return_value=None):
|
||||
manager.start('p',audio());self.assertEqual(done(manager)['state'],'complete');pid=manager.process.pid
|
||||
manager.start('p',audio());self.assertEqual(done(manager)['state'],'complete');self.assertEqual(manager.process.pid,pid)
|
||||
self.assertTrue(manager.status()['loaded'])
|
||||
finally:manager.stop()
|
||||
self.assertFalse(manager.status()['loaded'])
|
||||
|
||||
if __name__=='__main__':unittest.main()
|
||||
Reference in New Issue
Block a user