File size: 4,507 Bytes
31dc8dc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 | from __future__ import annotations
from typing import Any, Callable
from diffulex.logger import get_logger
from diffulex.server.protocol import (
ServingCommand,
ServingEvent,
ServingShutdown,
serving_command_from_dict,
serving_event_to_dict,
)
from diffulex.server.zmq_queue import ZmqPullQueue, ZmqPushQueue
logger = get_logger(__name__)
def default_engine_factory(model: str, **engine_kwargs):
from diffulex import strategy as _strategy # noqa: F401
from diffulex.engine.engine import DiffulexEngine
return DiffulexEngine(model, **engine_kwargs)
class SyncBackendWorker:
def __init__(
self,
*,
model: str,
engine_kwargs: dict[str, Any] | None,
recv_frontend,
send_frontend,
engine_factory: Callable[..., Any] | None = None,
ready_queue=None,
) -> None:
self.model = model
self.engine_kwargs = dict(engine_kwargs or {})
self.recv_frontend = recv_frontend
self.send_frontend = send_frontend
self.engine_factory = engine_factory or default_engine_factory
self.ready_queue = ready_queue
self.engine = None
self.shutdown_requested = False
@classmethod
def from_zmq(
cls,
*,
model: str,
engine_kwargs: dict[str, Any] | None,
command_addr: str,
event_addr: str,
ready_queue=None,
) -> "SyncBackendWorker":
return cls(
model=model,
engine_kwargs=engine_kwargs,
recv_frontend=ZmqPullQueue(command_addr, create=False, decoder=serving_command_from_dict),
send_frontend=ZmqPushQueue(event_addr, create=False, encoder=serving_event_to_dict),
ready_queue=ready_queue,
)
def init_engine(self) -> None:
if self.engine is not None:
return
self.engine = self.engine_factory(self.model, **self.engine_kwargs)
if self.ready_queue is not None:
self.ready_queue.put("SyncBackendWorker is ready")
def run_forever(self) -> None:
self.init_engine()
try:
while not self.shutdown_requested:
self.normal_loop()
finally:
self.shutdown()
def normal_loop(self) -> None:
assert self.engine is not None
commands = self.receive_commands(blocking=self.engine.is_finished())
commands = self.process_input_commands(commands)
if commands or not self.engine.is_finished():
events = self.engine.run_serving_tick(commands)
self.send_result(events)
def receive_commands(self, *, blocking: bool) -> list[ServingCommand]:
commands: list[ServingCommand] = []
if blocking:
self.run_when_idle()
commands.append(self.recv_frontend.get())
while not self.recv_frontend.empty():
commands.append(self.recv_frontend.get())
return commands
def process_input_commands(self, commands: list[ServingCommand]) -> list[ServingCommand]:
input_commands: list[ServingCommand] = []
for command in commands:
if isinstance(command, ServingShutdown):
self.shutdown_requested = True
else:
input_commands.append(command)
return input_commands
def send_result(self, events: list[ServingEvent]) -> None:
for event in events:
self.send_frontend.put(event)
def run_when_idle(self) -> None:
logger.info("SyncBackendWorker is idle, waiting for new requests...")
def shutdown(self) -> None:
if self.engine is not None:
try:
self.engine.exit()
finally:
self.engine = None
for queue in (self.recv_frontend, self.send_frontend):
stop = getattr(queue, "stop", None)
if stop is not None:
stop()
def run_sync_backend_worker(
*,
model: str,
engine_kwargs: dict[str, Any] | None,
command_addr: str,
event_addr: str,
ready_queue=None,
) -> None:
try:
worker = SyncBackendWorker.from_zmq(
model=model,
engine_kwargs=engine_kwargs,
command_addr=command_addr,
event_addr=event_addr,
ready_queue=ready_queue,
)
worker.run_forever()
except Exception as exc:
if ready_queue is not None:
ready_queue.put({"error": repr(exc)})
raise
|