import json from api import ComposerAPI, GenerationRequest from data import midi2wavtool, wavtool2midi from bottle import Bottle, run, request, response import threading from pathlib import Path from loader import load_by_path from util import configure_logging, stopwatch from warmup import warmup # import gil_load import sys def serve(path): path = Path(path) with stopwatch("startup"): api = ComposerAPI(*load_by_path(path)) warmup(api) app = Bottle() netlock = threading.Lock() @app.hook("after_request") def cors(): response.headers["Access-Control-Allow-Origin"] = "*" response.headers[ "Access-Control-Allow-Methods" ] = "PUT, GET, POST, DELETE, OPTIONS" response.headers[ "Access-Control-Allow-Headers" ] = "Origin, Accept, Content-Type, X-Requested-With, X-CSRF-Token" @app.route("/continue", method=["OPTIONS", "POST"]) def generate(): if request.method == "OPTIONS": return {} obj = request.json accompany = obj.get("accompany", None) if accompany is not None: accompany = wavtool2midi({"notes": accompany}) text_prompt = obj.get("textPrompt", None) with netlock: midi_out = api.generate( wavtool2midi(obj), [GenerationRequest.from_json(x) for x in obj["requests"]], accompany=accompany, text_prompt=text_prompt, ) response.content_type = "application/json" return json.dumps([midi2wavtool(x) for x in midi_out]) run(app, host="0.0.0.0", port=8081, debug=True) if __name__ == "__main__": import sys configure_logging() serve(sys.argv[1])