Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[project]
name = "ezmsg"
version = "3.10.0b4"
version = "3.10.0b5"
description = "A simple DAG-based computation model"
authors = [
{ name = "Griffin Milsap", email = "griffin.milsap@gmail.com" },
Expand Down
205 changes: 205 additions & 0 deletions src/ezmsg/core/axiselision.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,205 @@
"""
Wire elision of coordinate axes a receiver already holds.

An :class:`~ezmsg.util.messages.axisarray.AxisArray` stream sends the same
non-stream coordinate axes (channel labels, positions, ...) in every message,
often more bytes than the data itself. Within a process that costs nothing --
messages share the axis objects -- but across one, every message serializes,
copies and unpickles them again, and the receiver gets a new axis object per
message, so consumers can only compare it by value.

Here the publisher sends each such axis in full once (an :class:`_AxisDef`:
token plus axis) and thereafter as a 16-byte :class:`_AxisRef`. The receiving
channel keeps the axes it was given, copied out of the transport's memory, and
substitutes them for the references, so every message of a stream shares one
axis object on the far side too.

The token is derived from :attr:`CoordinateAxis.fingerprint` (a 64-bit content
digest), so equal axes share a token no matter which object carried them.

Only the axes of a top-level ``AxisArray`` message are elided, and never its
stream axis, whose values change every message. Everything else pickles as
before. Set ``EZMSG_DISABLE_AXIS_ELISION`` to turn the feature off.
"""

import hashlib
import os
import pickle
import typing

from collections import OrderedDict

ELISION_ENABLED = "EZMSG_DISABLE_AXIS_ELISION" not in os.environ

# Publisher: forget what was announced (forcing definitions again) past this
# many distinct axes. Receiver: keep at most this many.
MAX_ANNOUNCED = 256
MAX_TABLE = 256

_AxisArray: typing.Any = None
_CoordinateAxis: typing.Any = None
# Stream dimension assumed for a message that does not declare one, so that a
# per-message axis by that name is never elided; the shared definition is
# ezmsg.util.messages.axisarray.DEFAULT_STREAM_DIM.
_default_stream_dim: str | None = None


def _types() -> tuple[typing.Any, typing.Any]:
# Imported lazily: ezmsg.util.messages imports ezmsg.core.
global _AxisArray, _CoordinateAxis, _default_stream_dim
if _AxisArray is None:
from ..util.messages.axisarray import AxisArray, CoordinateAxis, DEFAULT_STREAM_DIM

_AxisArray, _CoordinateAxis, _default_stream_dim = AxisArray, CoordinateAxis, DEFAULT_STREAM_DIM
return _AxisArray, _CoordinateAxis


class MissingAxis(Exception):
"""A message referenced an axis this receiver does not hold."""


class _AxisRef:
"""Stands in, on the wire, for an axis the receiver already holds."""

__slots__ = ("token",)

def __init__(self, token: bytes) -> None:
self.token = token

def __reduce__(self):
return (_AxisRef, (self.token,))


class _AxisDef:
"""Carries an axis on the wire along with the token later messages will use."""

__slots__ = ("token", "axis")

def __init__(self, token: bytes, axis: typing.Any) -> None:
self.token = token
self.axis = axis

def __reduce__(self):
return (_AxisDef, (self.token, self.axis))


def wire_token(axis: typing.Any) -> bytes | None:
"""The axis's token, computed once per axis object; None if it has no
fingerprint (contents that cannot be digested are never elided)."""
d = axis.__dict__
token = d.get("_wire_token")
if token is None:
fp = axis.fingerprint
if fp is None:
return None
# The fingerprint holds a numpy dtype, slow to unpickle; a fixed-size
# digest of it is what travels.
token = d["_wire_token"] = hashlib.blake2b(pickle.dumps(fp, protocol=5), digest_size=16).digest()
return token


def _stream_dim(d: dict) -> str | None:
stream = d.get("stream_dim")
if stream is None and _default_stream_dim in d["dims"]:
stream = _default_stream_dim
return stream


class AxisElision:
"""Publisher side: which axes have been announced to the current channels."""

__slots__ = ("announced",)

def __init__(self) -> None:
self.announced: set[bytes] = set()

def reset(self) -> None:
"""Send every axis in full again (a channel joined, or asked)."""
self.announced.clear()

def wire(self, obj: typing.Any) -> typing.Any:
"""``obj``, or a shallow stand-in whose eligible axes are elided."""
AxisArray, CoordinateAxis = _types()
if not isinstance(obj, AxisArray):
return obj
d = obj.__dict__
axes = d["axes"]
stream = _stream_dim(d)
new_axes = None
for dim, axis in axes.items():
if dim == stream or type(axis) is not CoordinateAxis:
continue
token = wire_token(axis)
if token is None:
continue
if new_axes is None:
new_axes = dict(axes)
if token in self.announced:
new_axes[dim] = _AxisRef(token)
else:
if len(self.announced) >= MAX_ANNOUNCED:
self.announced.clear()
self.announced.add(token)
new_axes[dim] = _AxisDef(token, axis)
if new_axes is None:
return obj
# A bare copy of the message (no __init__, so no validation) differing
# only in its axes dict; the caller's message is untouched.
wired = object.__new__(type(obj))
wired.__dict__.update(d)
wired.__dict__["axes"] = new_axes
return wired


def _owned(axis: typing.Any) -> typing.Any:
"""A copy of ``axis`` whose data owns its memory (it may arrive as a view
into a channel's shared memory, which must not be pinned or read after the
slot is reused), keeping its cached fingerprint and token."""
import numpy as np

data = axis.data
if isinstance(data, np.ndarray) and data.flags.owndata:
return axis
out = object.__new__(type(axis))
out.__dict__.update(axis.__dict__)
out.__dict__["data"] = np.array(data, copy=True)
return out


class AxisTable:
"""Receiver side: the axes this channel has been given, by token."""

__slots__ = ("_axes",)

def __init__(self) -> None:
self._axes: "OrderedDict[bytes, typing.Any]" = OrderedDict()

def __len__(self) -> int:
return len(self._axes)

def resolve(self, obj: typing.Any) -> typing.Any:
"""Replace an ``AxisArray``'s elided axes in place; record new ones.

:raises MissingAxis: if a reference names an axis not held here.
"""
AxisArray, _ = _types()
if not isinstance(obj, AxisArray):
return obj
axes = obj.__dict__["axes"]
for dim, axis in axes.items():
kind = type(axis)
if kind is _AxisRef:
try:
axes[dim] = self._axes[axis.token]
except KeyError:
raise MissingAxis(axis.token.hex()) from None
elif kind is _AxisDef:
held = self._axes.get(axis.token)
if held is None:
held = self._axes[axis.token] = _owned(axis.axis)
if len(self._axes) > MAX_TABLE:
self._axes.popitem(last=False)
else:
self._axes.move_to_end(axis.token)
axes[dim] = held
return obj
4 changes: 2 additions & 2 deletions src/ezmsg/core/messagecache.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,7 +78,7 @@ def put_local(self, obj: typing.Any, msg_id: int) -> None:
)
)

def put_from_mem(self, mem: memoryview) -> None:
def put_from_mem(self, mem: memoryview, axis_table: typing.Any = None) -> None:
"""
Reconstitute a message in mem and keep it in cache, releasing and
overwriting the existing slot in cache.
Expand All @@ -89,7 +89,7 @@ def put_from_mem(self, mem: memoryview) -> None:
:type from_mem: memoryview
:raises UninitializedMemory: If mem buffer is not properly initialized.
"""
ctx = MessageMarshal.obj_from_mem(mem)
ctx = MessageMarshal.obj_from_mem(mem, axis_table)
self._put(
CacheEntry(
object=ctx.__enter__(),
Expand Down
42 changes: 39 additions & 3 deletions src/ezmsg/core/messagechannel.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from uuid import UUID
from contextlib import contextmanager, suppress

from .axiselision import AxisTable, MissingAxis, ELISION_ENABLED
from .shm import SHMContext
from .messagemarshal import MessageMarshal, UninitializedMemory
from .backpressure import Backpressure
Expand Down Expand Up @@ -149,7 +150,8 @@ async def _reattach_shm(self, shm_name: str, frame_end: int, msg_id: int) -> Non
preserved = chan._snapshot_cached_messages()
chan.cache.clear()
for preserved_msg in preserved:
chan.cache.put_from_mem(preserved_msg)
with suppress(MissingAxis):
chan.cache.put_from_mem(preserved_msg, chan._axis_table)

if chan.shm is not None:
old_shm = chan.shm
Expand Down Expand Up @@ -245,6 +247,9 @@ def __init__(
self._graph_address = graph_address
self._local_backpressure = None
self._channel_kind = ProfileChannelType.UNKNOWN
# Axes the publisher sent in full, for resolving its elided references.
self._axis_table = AxisTable()
self._axes_requested = False

@classmethod
async def create(
Expand Down Expand Up @@ -311,6 +316,10 @@ async def create(
if num_buffers <= 0:
proto.close()
raise ValueError("publisher reports invalid num_buffers")
if ELISION_ENABLED:
# Tell the publisher we can resolve elided axes. An older publisher's
# read loop ignores the byte and keeps sending axes in full.
proto.write(Command.ELIDE_OK.value)

chan = cls(UUID(id_str), pub_id, num_buffers, shm, graph_address, _guard=cls._SENTINEL)
chan.topic = topic
Expand Down Expand Up @@ -403,7 +412,11 @@ def _deliver_from_shm(self, msg_id: int) -> None:
self._release_backpressure(msg_id, self.id)
return

self.cache.put_from_mem(shm_buf)
try:
self.cache.put_from_mem(shm_buf, self._axis_table)
except MissingAxis as exc:
self._missing_axis(msg_id, exc)
return
self._set_channel_kind(ProfileChannelType.SHM)
self._finish_delivery(msg_id)

Expand All @@ -414,11 +427,34 @@ def _deliver_from_tcp(self, msg_id: int, obj_bytes: bytes) -> None:
Called inline from :meth:`ChannelProtocol.frames_available`.
"""
assert MessageMarshal.msg_id(obj_bytes) == msg_id
self.cache.put_from_mem(memoryview(obj_bytes).toreadonly())
try:
self.cache.put_from_mem(memoryview(obj_bytes).toreadonly(), self._axis_table)
except MissingAxis as exc:
self._missing_axis(msg_id, exc)
return
self._set_channel_kind(ProfileChannelType.TCP)
self._finish_delivery(msg_id)

def _missing_axis(self, msg_id: int, exc: MissingAxis) -> None:
"""A message referenced an axis we do not hold (we missed or evicted
its definition). Drop it, as the stale-SHM path does, and ask the
publisher to send its axes in full again -- once per episode, since
several messages already in flight may reference it too."""
logger.warning(
"Channel %s dropping message %s from publisher %s: unknown axis %s",
self.id,
msg_id,
self.pub_id,
exc,
)
if not self._axes_requested:
self._axes_requested = True
self._proto.write(Command.AXIS_RESEND.value)
self._release_backpressure(msg_id, self.id)

def _finish_delivery(self, msg_id: int) -> None:
# A message resolved, so any axes we asked for have arrived.
self._axes_requested = False
if not self._notify_clients(msg_id):
# Nobody is listening; need to ack!
self.cache.release(msg_id)
Expand Down
13 changes: 11 additions & 2 deletions src/ezmsg/core/messagemarshal.py
Original file line number Diff line number Diff line change
Expand Up @@ -106,7 +106,7 @@ def msg_id(cls, raw: memoryview | bytes) -> int:

@classmethod
@contextmanager
def obj_from_mem(cls, mem: memoryview) -> Generator[Any, None, None]:
def obj_from_mem(cls, mem: memoryview, axis_table: Any = None) -> Generator[Any, None, None]:
"""
Deserialize an object from a memory buffer.

Expand All @@ -115,9 +115,12 @@ def obj_from_mem(cls, mem: memoryview) -> Generator[Any, None, None]:

:param mem: Memory buffer containing serialized object.
:type mem: memoryview
:param axis_table: The receiving channel's
:class:`~ezmsg.core.axiselision.AxisTable`, to resolve elided axes.
:return: Context manager yielding the deserialized object.
:rtype: Generator[Any, None, None]
:raises UninitializedMemory: If memory buffer is not properly initialized.
:raises MissingAxis: If the message references an axis not in ``axis_table``.
"""
cls._assert_initialized(mem)

Expand All @@ -138,6 +141,8 @@ def obj_from_mem(cls, mem: memoryview) -> Generator[Any, None, None]:
sidx += bsz

obj = cls.load(buffers)
if axis_table is not None:
axis_table.resolve(obj)

try:
yield obj
Expand All @@ -149,7 +154,7 @@ def obj_from_mem(cls, mem: memoryview) -> Generator[Any, None, None]:
@classmethod
@contextmanager
def serialize(
cls, msg_id: int, obj: Any
cls, msg_id: int, obj: Any, elision: Any = None
) -> Generator[tuple[int, bytes, list[memoryview]], None, None]:
"""
Serialize an object for network transmission.
Expand All @@ -161,9 +166,13 @@ def serialize(
:type msg_id: int
:param obj: Object to serialize.
:type obj: Any
:param elision: The publisher's :class:`~ezmsg.core.axiselision.AxisElision`,
when every receiver can resolve elided axes.
:return: Context manager yielding (total_size, header, buffers) tuple.
:rtype: Generator[tuple[int, bytes, list[memoryview]], None, None]
"""
if elision is not None:
obj = elision.wire(obj)
buffers = cls.dump(obj)
header = uint64_to_bytes(len(buffers))
buf_lengths = [len(buf) for buf in buffers]
Expand Down
7 changes: 7 additions & 0 deletions src/ezmsg/core/netprotocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -347,6 +347,13 @@ def _generate_next_value_(name, start, count, last_values) -> bytes:
PROCESS_ROUTE_RESPONSE = enum.auto()
ERROR = enum.auto()

# Channel -> Publisher: axis elision (appended, so no existing value moves).
# A channel that can resolve elided axes says so once after connecting; an
# older publisher's read loop ignores the byte.
ELIDE_OK = enum.auto()
# A channel received a reference to an axis it does not hold.
AXIS_RESEND = enum.auto()


def create_socket(
host: str | None = None,
Expand Down
Loading