Skip to content

Commit 46806ae

Browse files
fix(server): make uvicorn workers>1 work (factory + spawn for CUDA) (#1141)
* fix: pass cli args to uv workers via env vars * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * fix: multi-worker mode with spawns instead of default forks for CUDA --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
1 parent ba8bdfb commit 46806ae

1 file changed

Lines changed: 31 additions & 7 deletions

File tree

tools/api_server.py

Lines changed: 31 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,8 @@
1+
import json
2+
import multiprocessing
3+
import os
14
import re
5+
from argparse import Namespace
26
from threading import Lock
37

48
import pyrootutils
@@ -25,10 +29,12 @@
2529
from tools.server.model_manager import ModelManager
2630
from tools.server.views import routes
2731

32+
ENV_ARGS_KEY = "FISH_API_SERVER_ARGS"
33+
2834

2935
class API(ExceptionHandler):
30-
def __init__(self):
31-
self.args = parse_args()
36+
def __init__(self, args: Namespace | None = None):
37+
self.args = args or parse_args()
3238

3339
def api_auth(endpoint):
3440
async def verify(token: Annotated[str, Depends(bearer_auth)]):
@@ -93,6 +99,19 @@ async def initialize_app(self, app: Kui):
9399
logger.info(f"Startup done, listening server at http://{self.args.listen}")
94100

95101

102+
def create_app():
103+
args_env = os.environ.get(ENV_ARGS_KEY)
104+
args = None
105+
106+
if args_env:
107+
try:
108+
args = Namespace(**json.loads(args_env))
109+
except Exception as exc:
110+
logger.warning(f"Failed to load args from {ENV_ARGS_KEY}: {exc}")
111+
112+
return API(args=args).app
113+
114+
96115
# Each worker process created by Uvicorn has its own memory space,
97116
# meaning that models and variables are not shared between processes.
98117
# Therefore, any variables (like `llama_queue` or `decoder_model`)
@@ -103,19 +122,24 @@ async def initialize_app(self, app: Kui):
103122
# Instead, it's better to use multiprocessing or independent models per thread.
104123

105124
if __name__ == "__main__":
106-
api = API()
125+
126+
multiprocessing.set_start_method("spawn", force=True)
127+
128+
args = parse_args()
129+
os.environ[ENV_ARGS_KEY] = json.dumps(vars(args))
107130

108131
# IPv6 address format is [xxxx:xxxx::xxxx]:port
109-
match = re.search(r"\[([^\]]+)\]:(\d+)$", api.args.listen)
132+
match = re.search(r"\[([^\]]+)\]:(\d+)$", args.listen)
110133
if match:
111134
host, port = match.groups() # IPv6
112135
else:
113-
host, port = api.args.listen.split(":") # IPv4
136+
host, port = args.listen.split(":") # IPv4
114137

115138
uvicorn.run(
116-
api.app,
139+
"tools.api_server:create_app",
117140
host=host,
118141
port=int(port),
119-
workers=api.args.workers,
142+
workers=args.workers,
120143
log_level="info",
144+
factory=True,
121145
)

0 commit comments

Comments
 (0)