sparx.serve
Serving streaming spiking models: many sessions, each with its own neuron state, in one batched program.
server = StreamServer(model, variables, slots=8, frame=10, sample_shape=(40,))session = server.open()future = server.submit(session, chunk) # [10, 40]: one frame of this session's streamserver.run() # one batched step per round of pending framesoutputs = future.result() # [10, ...]A run that dew.pipeline reloads serves as StreamServer(classifier.model, classifier.variables, call=classifier.call, ...); its frames are the
encoder’s output, time-major.
A spiking model that streams (the state collection, the guide’s “Streaming”)
carries its neurons’ state from one call to the next. A server keeps that
state for slots sessions as the rows of one resident batch: each step
runs one frame of every session that has one waiting, in one jitted call,
and a session without a frame keeps its state exactly (its row is not
advanced, not advanced on zeros). A session’s outputs over its frames are
then the outputs of one call over its whole stream. Closing a session
frees its row.
A session’s first frame runs from rest, the state every layer starts from
when the state collection holds nothing. The server never builds a rest
state itself: a step whose rows include a session’s first frame also runs
the model without carried state and takes those rows from that run. Any
model whose layers follow the state convention serves this way, whatever
it does between them, and only a step that starts a session pays for a
second pass.
This is dew’s slot-scheduling Server for text (a resident KV cache
refilled from a queue) applied to neuron state. dew’s own server overlaps
dispatch with token draws, which synchronous frames do not need, and making
it generic was declined (AshishKumar4/dew#30).
Contents
Section titled “Contents”| Name | |
|---|---|
StreamServer | Sessions of a streaming model, slots at a time, advanced a frame of steps per round. |
StreamServer
Section titled “StreamServer”class StreamServerSessions of a streaming model, slots at a time, advanced a frame of steps per round.
variables holds every collection the model reads (params, and
batch_stats for a model with batch norm); a state collection in it
is ignored, since each session starts at rest. The model runs as a
trained classifier does (sparx.tasks.bound_call): with train=False
when it takes train, and the keyword arguments call (a loaded
classifier’s call, its schedules’ last values). Construction runs the
model once on abstract inputs, so a model that cannot stream fails here
with the reason.
A future cancelled before its frame runs drops that frame, and the session runs its next one. A step that raises sets the exception on the futures of every frame it took, leaves every session’s state as it was, and raises it.
| Field | Type | Default |
|---|---|---|
state | Variables | None | None |
sessions | dict[int, int] | {} |
fresh | set[int] | set() |
pending | dict[int, collections.deque[tuple[np.ndarray, Future]]] | {} |
StreamServer.open
Section titled “StreamServer.open”def open() -> intStart a session at rest in a free slot; returns its id.
StreamServer.close
Section titled “StreamServer.close”def close(session: int) -> NoneEnd a session; its slot takes the next one. Frames still waiting are cancelled.
StreamServer.submit
Section titled “StreamServer.submit”def submit(session: int, frame: np.ndarray | jax.Array) -> FutureQueue one frame [frame, *sample_shape] of session’s stream; the future holds its outputs.
StreamServer.step
Section titled “StreamServer.step”def step() -> intRun the oldest waiting frame of every session that has one; returns how many ran.
StreamServer.run
Section titled “StreamServer.run”def run() -> NoneStep until no frame waits.