Skip to content
GitHub

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 stream
server.run() # one batched step per round of pending frames
outputs = 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).

Name
StreamServerSessions of a streaming model, slots at a time, advanced a frame of steps per round.
class StreamServer

sparx.serve on GitHub

Sessions 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.

FieldTypeDefault
stateVariables | NoneNone
sessionsdict[int, int]{}
freshset[int]set()
pendingdict[int, collections.deque[tuple[np.ndarray, Future]]]{}
def open() -> int

Start a session at rest in a free slot; returns its id.

def close(session: int) -> None

End a session; its slot takes the next one. Frames still waiting are cancelled.

def submit(session: int, frame: np.ndarray | jax.Array) -> Future

Queue one frame [frame, *sample_shape] of session’s stream; the future holds its outputs.

def step() -> int

Run the oldest waiting frame of every session that has one; returns how many ran.

def run() -> None

Step until no frame waits.