MATE-libraries 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,5 @@
1
+ """MATE libraries: the network library and the standard operators library, for MATE modules."""
2
+
3
+ from importlib.metadata import version
4
+
5
+ __version__ = version("MATE-libraries")
@@ -0,0 +1,67 @@
1
+ """The network library: compile and propagate batches of ``Network`` organs.
2
+
3
+ Three objects are never fused: the organ (frozen arrays, one per individual), the compiled
4
+ network (immutable, batched, shared between threads, never serialised) and the activation
5
+ state (the node values of individual x episode rows, mutable and serialised by copy). The
6
+ caller chooses what changes the result (propagation mode, refresh rate, activation functions,
7
+ aggregation, fate of unfilled cells, compute dtype); the library chooses only what does not
8
+ (batching technique, shortcuts derived from the graph). Each individual's results are
9
+ bit-identical whether it is compiled alone or in any batch, under any technique.
10
+ ``from_networkx`` builds an organ from a networkx graph, for hand-designed networks.
11
+
12
+ Attributes
13
+ ----------
14
+ version : int
15
+ What the library computes; changes whenever a definition, a default, a convergence rule
16
+ or the set of accepted inputs changes (rules V1, V4 of the libraries specification). A
17
+ protocol may hold it fixed under the key ``networks.version``.
18
+ TECHNIQUES : tuple of str
19
+ The batching techniques: ``"dense"``, ``"padded"``, ``"block_diagonal"``.
20
+ ACTIVATIONS : Mapping[str, Activation]
21
+ Read-only catalogue of activation functions: ``id``, ``tanh``, ``relu``, ``sigmoid``,
22
+ ``neat_sigmoid``, ``sin``, ``gauss``, ``abs``, ``step``.
23
+ AGGREGATIONS : Mapping[str, Aggregation]
24
+ Read-only catalogue of aggregations: ``sum`` (the default), ``product``, ``max``,
25
+ ``min``, ``mean``.
26
+
27
+ Notes
28
+ -----
29
+ Section 4 of the libraries specification. The device path is deferred in MATE-libraries
30
+ 1.0.0: ``Device("opencl")`` is refused, any other non-CPU device runs on the CPU, and
31
+ ``CompiledNetwork.kernels()`` raises.
32
+
33
+ Examples
34
+ --------
35
+ >>> from MATE_libraries import networks as nets
36
+ >>> compiled = nets.compile(organs, spec, unfilled={"activation": "id"},
37
+ ... device=ctx.device) # doctest: +SKIP
38
+ >>> act = compiled.state(10) # doctest: +SKIP
39
+ >>> y = compiled.step(act, x[t], refresh_rate=1) # doctest: +SKIP
40
+ >>> res = compiled.settle(u, max_steps=50, tol=1e-6) # doctest: +SKIP
41
+ """
42
+
43
+ from typing import Final
44
+
45
+ from ._catalogue import ACTIVATIONS, AGGREGATIONS, VERSION, Activation, Aggregation
46
+ from ._compiled import TECHNIQUES, CompileDiagnostics, CompiledNetwork, compile, validate
47
+ from ._networkx import from_networkx
48
+ from ._state import ActivationState, IndividualActivation, SettleResult
49
+
50
+ version: Final[int] = VERSION
51
+
52
+ __all__ = [
53
+ "version",
54
+ "TECHNIQUES",
55
+ "Activation",
56
+ "Aggregation",
57
+ "ACTIVATIONS",
58
+ "AGGREGATIONS",
59
+ "compile",
60
+ "validate",
61
+ "from_networkx",
62
+ "CompileDiagnostics",
63
+ "CompiledNetwork",
64
+ "ActivationState",
65
+ "IndividualActivation",
66
+ "SettleResult",
67
+ ]
@@ -0,0 +1,282 @@
1
+ """Catalogues of activation functions and aggregations, and name resolution ([L] §4.5)."""
2
+
3
+ from collections.abc import Callable, Mapping
4
+ from dataclasses import dataclass, field
5
+ from types import MappingProxyType
6
+ from typing import Any, Final
7
+
8
+ import numpy as np
9
+
10
+ VERSION: Final[int] = 2
11
+ """What the network library computes (rules V1, V4); exposed as ``networks.version``."""
12
+
13
+
14
+ def _kernels(value: Mapping[str, str], where: str) -> Mapping[str, str]:
15
+ """Return a read-only copy of a backend-to-kernel-source mapping."""
16
+ if not isinstance(value, Mapping):
17
+ raise TypeError(f"{where}.kernels: expected a mapping, got {type(value).__name__}")
18
+ for backend, source in value.items():
19
+ if not isinstance(backend, str) or not isinstance(source, str):
20
+ raise TypeError(f"{where}.kernels: keys and values are str (backend -> source)")
21
+ return MappingProxyType(dict(value))
22
+
23
+
24
+ @dataclass(frozen=True)
25
+ class Activation:
26
+ """An activation function of the network library.
27
+
28
+ Parameters
29
+ ----------
30
+ function : callable
31
+ Elementwise, pure and dtype-preserving map of an array; its result for an element
32
+ must not depend on the element's position, the array length, alignment or strides.
33
+ kernels : mapping of str to str, default {}
34
+ Device kernel source per backend name; without one, the device path is lost. The
35
+ ``"opencl"`` source defines ``real activate(real x)``, where ``real`` is the compute
36
+ type (``float`` or ``double``) and ``MATE_REAL(1.0)`` a literal of that type, with
37
+ the operations of ``function`` in the same order.
38
+
39
+ Raises
40
+ ------
41
+ TypeError
42
+ If ``function`` is not callable or ``kernels`` is not a mapping of str to str.
43
+
44
+ Notes
45
+ -----
46
+ Rules K2, K5 of the libraries specification. Pickles through its constructor, so that a
47
+ module keeping it can be sent to a worker process.
48
+ """
49
+
50
+ function: Callable[[np.ndarray], np.ndarray]
51
+ kernels: Mapping[str, str] = field(default_factory=dict)
52
+
53
+ def __post_init__(self) -> None:
54
+ if not callable(self.function):
55
+ raise TypeError(
56
+ f"Activation.function: expected a callable, got {type(self.function).__name__}"
57
+ )
58
+ object.__setattr__(self, "kernels", _kernels(self.kernels, "Activation"))
59
+
60
+ def __reduce__(self) -> tuple[Any, ...]:
61
+ return (Activation, (self.function, dict(self.kernels)))
62
+
63
+
64
+ @dataclass(frozen=True)
65
+ class Aggregation:
66
+ """How a node combines the weighted terms of its present incoming edges.
67
+
68
+ Parameters
69
+ ----------
70
+ combine : callable
71
+ ``(accumulator, term) -> accumulator``, elementwise; the fold starts from the first
72
+ term and follows the canonical edge order.
73
+ empty : float
74
+ Value of a node without any present incoming edge.
75
+ finalize : callable, optional
76
+ ``(accumulator, count) -> value``, applied after the fold.
77
+ kernels : mapping of str to str, default {}
78
+ Device kernel source per backend name. The ``"opencl"`` source defines
79
+ ``real combine(real acc, real term)`` and, when ``finalize`` is given,
80
+ ``real finalize(real acc, int count)``, on the compute type ``real`` (see
81
+ :class:`Activation`).
82
+
83
+ Raises
84
+ ------
85
+ TypeError
86
+ If ``combine`` or ``finalize`` is not callable, ``empty`` is not a real number, or
87
+ ``kernels`` is not a mapping of str to str.
88
+
89
+ Notes
90
+ -----
91
+ Rules K2, K4, K5 and S9a of the libraries specification. Pickles through its constructor,
92
+ so that a module keeping it can be sent to a worker process.
93
+ """
94
+
95
+ combine: Callable[[np.ndarray, np.ndarray], np.ndarray]
96
+ empty: float
97
+ finalize: Callable[[np.ndarray, np.ndarray], np.ndarray] | None = None
98
+ kernels: Mapping[str, str] = field(default_factory=dict)
99
+
100
+ def __post_init__(self) -> None:
101
+ if not callable(self.combine):
102
+ raise TypeError(
103
+ f"Aggregation.combine: expected a callable, got {type(self.combine).__name__}"
104
+ )
105
+ if self.finalize is not None and not callable(self.finalize):
106
+ raise TypeError(
107
+ "Aggregation.finalize: expected a callable or None, got "
108
+ f"{type(self.finalize).__name__}"
109
+ )
110
+ if isinstance(self.empty, bool) or not isinstance(self.empty, (int, float, np.number)):
111
+ raise TypeError(
112
+ f"Aggregation.empty: expected a real number, got {type(self.empty).__name__}"
113
+ )
114
+ object.__setattr__(self, "empty", float(self.empty))
115
+ object.__setattr__(self, "kernels", _kernels(self.kernels, "Aggregation"))
116
+
117
+ def __reduce__(self) -> tuple[Any, ...]:
118
+ return (Aggregation, (self.combine, self.empty, self.finalize, dict(self.kernels)))
119
+
120
+
121
+ # Catalogue functions: built from numpy ufuncs evaluated in the array's own dtype, so that
122
+ # each element goes through one code path whatever its position (K5).
123
+
124
+
125
+ def _identity(x: np.ndarray) -> np.ndarray:
126
+ """x (a new array, so that callers never alias their input)."""
127
+ return np.positive(x)
128
+
129
+
130
+ def _tanh(x: np.ndarray) -> np.ndarray:
131
+ """tanh(x)."""
132
+ return np.tanh(x)
133
+
134
+
135
+ def _relu(x: np.ndarray) -> np.ndarray:
136
+ """max(x, 0)."""
137
+ return np.maximum(x, x.dtype.type(0))
138
+
139
+
140
+ def _sigmoid(x: np.ndarray) -> np.ndarray:
141
+ """1 / (1 + exp(-x))."""
142
+ one = x.dtype.type(1)
143
+ return one / (one + np.exp(-x))
144
+
145
+
146
+ def _neat_sigmoid(x: np.ndarray) -> np.ndarray:
147
+ """1 / (1 + exp(-4.9 x)), the steepened sigmoid of the original NEAT paper."""
148
+ one = x.dtype.type(1)
149
+ return one / (one + np.exp(x.dtype.type(-4.9) * x))
150
+
151
+
152
+ def _sin(x: np.ndarray) -> np.ndarray:
153
+ """sin(x)."""
154
+ return np.sin(x)
155
+
156
+
157
+ def _gauss(x: np.ndarray) -> np.ndarray:
158
+ """exp(-x²)."""
159
+ return np.exp(-(x * x))
160
+
161
+
162
+ def _abs(x: np.ndarray) -> np.ndarray:
163
+ """|x|."""
164
+ return np.abs(x)
165
+
166
+
167
+ def _step(x: np.ndarray) -> np.ndarray:
168
+ """1 if x > 0 else 0."""
169
+ return (x > 0).astype(x.dtype)
170
+
171
+
172
+ def _mean(acc: np.ndarray, count: np.ndarray) -> np.ndarray:
173
+ """Sum divided by the term count (0 where there is no term)."""
174
+ count = np.asarray(count).astype(acc.dtype)
175
+ return np.divide(acc, count, out=np.zeros_like(acc), where=count > 0)
176
+
177
+
178
+ def _opencl(body: str, signature: str = "real activate(real x)") -> dict[str, str]:
179
+ """The ``kernels`` mapping of a catalogue entry: one OpenCL function."""
180
+ return {"opencl": f"inline {signature} {{ {body} }}"}
181
+
182
+
183
+ # The OpenCL sources repeat the numpy expressions operation by operation. numpy's maximum and
184
+ # minimum propagate a NaN of either side and, on equal operands (+0 and -0 included), return
185
+ # the second one, as the x86 vector instructions they run on do.
186
+ ACTIVATIONS: Final[Mapping[str, Activation]] = MappingProxyType(
187
+ {
188
+ "id": Activation(_identity, _opencl("return x;")),
189
+ "tanh": Activation(_tanh, _opencl("return tanh(x);")),
190
+ "relu": Activation(_relu, _opencl(
191
+ "return (isnan(x) || x > MATE_REAL(0.0)) ? x : MATE_REAL(0.0);")),
192
+ "sigmoid": Activation(_sigmoid, _opencl(
193
+ "return MATE_REAL(1.0) / (MATE_REAL(1.0) + exp(-x));")),
194
+ "neat_sigmoid": Activation(_neat_sigmoid, _opencl(
195
+ "return MATE_REAL(1.0) / (MATE_REAL(1.0) + exp(-MATE_REAL(4.9) * x));")),
196
+ "sin": Activation(_sin, _opencl("return sin(x);")),
197
+ "gauss": Activation(_gauss, _opencl("return exp(-(x * x));")),
198
+ "abs": Activation(_abs, _opencl("return fabs(x);")),
199
+ "step": Activation(_step, _opencl(
200
+ "return x > MATE_REAL(0.0) ? MATE_REAL(1.0) : MATE_REAL(0.0);")),
201
+ }
202
+ )
203
+ """Read-only catalogue of activation functions, dated by ``VERSION`` (rule K3)."""
204
+
205
+ _COMBINE = "real combine(real acc, real term)"
206
+ AGGREGATIONS: Final[Mapping[str, Aggregation]] = MappingProxyType(
207
+ {
208
+ "sum": Aggregation(np.add, 0.0, kernels=_opencl("return acc + term;", _COMBINE)),
209
+ "product": Aggregation(np.multiply, 1.0,
210
+ kernels=_opencl("return acc * term;", _COMBINE)),
211
+ "max": Aggregation(np.maximum, 0.0, kernels=_opencl(
212
+ "return (isnan(acc) || acc > term) ? acc : term;", _COMBINE)),
213
+ "min": Aggregation(np.minimum, 0.0, kernels=_opencl(
214
+ "return (isnan(acc) || acc < term) ? acc : term;", _COMBINE)),
215
+ "mean": Aggregation(np.add, 0.0, finalize=_mean, kernels={"opencl": (
216
+ f"inline {_COMBINE} {{ return acc + term; }}\n"
217
+ "inline real finalize(real acc, int count) { return acc / (real)count; }")}),
218
+ }
219
+ )
220
+ """Read-only catalogue of aggregations, dated by ``VERSION`` (rule K4)."""
221
+
222
+ _KINDS: Final[Mapping[str, tuple[Mapping[str, object], type, str]]] = MappingProxyType(
223
+ {
224
+ "activation": (ACTIVATIONS, Activation, "activations"),
225
+ "aggregation": (AGGREGATIONS, Aggregation, "aggregations"),
226
+ }
227
+ )
228
+
229
+
230
+ def caller_definitions[T](
231
+ given: Mapping[str, T] | None, kind: str
232
+ ) -> Mapping[str, T]:
233
+ """Check the definitions a caller supplies for ``kind`` and return them (R2, K2, K6)."""
234
+ catalogue, cls, keyword = _KINDS[kind]
235
+ if given is None:
236
+ return MappingProxyType({})
237
+ if not isinstance(given, Mapping):
238
+ raise TypeError(f"{keyword}=: expected a mapping of str to {cls.__name__}, "
239
+ f"got {type(given).__name__}")
240
+ for name, definition in given.items():
241
+ if not isinstance(name, str):
242
+ raise TypeError(f"{keyword}=: names are str, got {type(name).__name__}")
243
+ if not isinstance(definition, cls):
244
+ raise TypeError(
245
+ f"{keyword}['{name}']: expected an {cls.__name__}, got "
246
+ f"{type(definition).__name__}"
247
+ )
248
+ if name in catalogue:
249
+ raise ValueError(
250
+ f"'{name}' is a catalogue name of the network library (version {VERSION}); "
251
+ "give your variant another name"
252
+ )
253
+ return MappingProxyType(dict(given))
254
+
255
+
256
+ def missing_names_message(kind: str, missing: Mapping[str, list[str]]) -> str:
257
+ """Build the K6 message for names of ``kind`` found nowhere.
258
+
259
+ ``missing`` maps each unknown name to the places it was read from, in order of appearance.
260
+ """
261
+ catalogue, cls, keyword = _KINDS[kind]
262
+ names = list(missing)
263
+ places: list[str] = []
264
+ for where in missing.values():
265
+ places.extend(w for w in where if w not in places)
266
+ quoted = ", ".join(f"'{n}'" for n in names)
267
+ supply = ", ".join(f"'{n}': {cls.__name__}(...)" for n in names)
268
+ noun = kind if len(names) == 1 else f"{kind}s"
269
+ pronoun = "it" if len(names) == 1 else "them"
270
+ return (
271
+ f"network library (version {VERSION}): unknown {noun} {quoted} in "
272
+ f"{' and '.join(places)}; the catalogue offers {list(catalogue)}; "
273
+ f"pass {keyword}={{{supply}}} to supply {pronoun}"
274
+ )
275
+
276
+
277
+ def lookup[T](name: str, kind: str, given: Mapping[str, T]) -> T | None:
278
+ """Resolve ``name``: the catalogue first, then the caller's definitions (K6)."""
279
+ catalogue = _KINDS[kind][0]
280
+ if name in catalogue:
281
+ return catalogue[name] # type: ignore[return-value]
282
+ return given.get(name)