graphspace 0.1.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.
- graphspace/__init__.py +16 -0
- graphspace/_version.py +1 -0
- graphspace/backends/__init__.py +28 -0
- graphspace/backends/numpy.py +127 -0
- graphspace/backends/python.py +184 -0
- graphspace/claims.py +28 -0
- graphspace/cli.py +24 -0
- graphspace/core.py +697 -0
- graphspace/failures.py +64 -0
- graphspace/integrations/__init__.py +23 -0
- graphspace/integrations/numpy.py +14 -0
- graphspace/integrations/torch.py +14 -0
- graphspace/py.typed +0 -0
- graphspace/uncertainty.py +34 -0
- graphspace-0.1.0.dist-info/METADATA +101 -0
- graphspace-0.1.0.dist-info/RECORD +20 -0
- graphspace-0.1.0.dist-info/WHEEL +5 -0
- graphspace-0.1.0.dist-info/entry_points.txt +2 -0
- graphspace-0.1.0.dist-info/licenses/LICENSE +202 -0
- graphspace-0.1.0.dist-info/top_level.txt +1 -0
graphspace/core.py
ADDED
|
@@ -0,0 +1,697 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from collections.abc import Mapping, Sequence
|
|
4
|
+
from dataclasses import asdict, dataclass
|
|
5
|
+
from datetime import datetime, timezone
|
|
6
|
+
from typing import Any
|
|
7
|
+
import hashlib
|
|
8
|
+
import json
|
|
9
|
+
import platform
|
|
10
|
+
import sys
|
|
11
|
+
import threading
|
|
12
|
+
|
|
13
|
+
from ._version import __version__
|
|
14
|
+
from .claims import Basis, Claim, find_claim
|
|
15
|
+
from .failures import ContractViolation, DTypeMismatch, GraphspaceError, ResourceLimitExceeded, ShapeMismatch, UnknownValue
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
DTYPE_BYTES = {"float32": 4, "float16": 2, "int32": 4, "int64": 8}
|
|
19
|
+
LAYOUTS = frozenset({"row_major"})
|
|
20
|
+
FLOAT_DTYPES = frozenset({"float32", "float16"})
|
|
21
|
+
INT_RANGES = {"int32": (-(2**31), 2**31 - 1), "int64": (-(2**63), 2**63 - 1)}
|
|
22
|
+
IN_PLACE = frozenset({"add", "multiply", "subtract", "divide", "relu", "scale", "softmax", "layer_norm"})
|
|
23
|
+
ROW_SCRATCH = frozenset({"softmax", "layer_norm"})
|
|
24
|
+
UNBUFFERED = frozenset({"reshape", "transpose", "matmul"})
|
|
25
|
+
UFUNC_BUFFER_ELEMENTS = 8192
|
|
26
|
+
|
|
27
|
+
Dims = Mapping[str, int]
|
|
28
|
+
PREPARED_CACHE_SIZE = 8
|
|
29
|
+
BOOKKEEPING_BY_VERSION = {
|
|
30
|
+
(3, 10): (12992, 3208, 424),
|
|
31
|
+
(3, 11): (14208, 3160, 312),
|
|
32
|
+
(3, 12): (7744, 2728, 96),
|
|
33
|
+
(3, 13): (7552, 2720, 112),
|
|
34
|
+
(3, 14): (6080, 2664, 0),
|
|
35
|
+
}
|
|
36
|
+
BOOKKEEPING_BASE_BYTES, BOOKKEEPING_INPUT_BYTES, BOOKKEEPING_NODE_BYTES = BOOKKEEPING_BY_VERSION.get(
|
|
37
|
+
sys.version_info[:2], tuple(max(column) for column in zip(*BOOKKEEPING_BY_VERSION.values())),
|
|
38
|
+
)
|
|
39
|
+
PYTHON_VERSION = f"{platform.python_implementation()} {platform.python_version()}"
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
@dataclass(frozen=True)
|
|
43
|
+
class TensorSpec:
|
|
44
|
+
shape: tuple[Any, ...]
|
|
45
|
+
dtype: str = "float32"
|
|
46
|
+
layout: str = "row_major"
|
|
47
|
+
role: str = "temporary"
|
|
48
|
+
|
|
49
|
+
def __post_init__(self) -> None:
|
|
50
|
+
if self.dtype not in DTYPE_BYTES:
|
|
51
|
+
raise DTypeMismatch(
|
|
52
|
+
f"unsupported dtype: {self.dtype}",
|
|
53
|
+
expected=sorted(DTYPE_BYTES), actual=self.dtype,
|
|
54
|
+
)
|
|
55
|
+
if self.layout not in LAYOUTS:
|
|
56
|
+
raise ContractViolation(f"unsupported layout: {self.layout}", expected=sorted(LAYOUTS), actual=self.layout)
|
|
57
|
+
if isinstance(self.shape, (str, bytes)) or not isinstance(self.shape, Sequence):
|
|
58
|
+
raise ShapeMismatch(f"invalid shape: {self.shape!r}", expected="sequence of dimensions", actual=self.shape)
|
|
59
|
+
object.__setattr__(self, "shape", tuple(self.shape))
|
|
60
|
+
for value in self.shape:
|
|
61
|
+
if isinstance(value, bool) or not isinstance(value, (int, str)):
|
|
62
|
+
raise ShapeMismatch(
|
|
63
|
+
f"invalid dimension {value!r} in shape {self.shape}",
|
|
64
|
+
expected="int or symbolic name", actual=value,
|
|
65
|
+
)
|
|
66
|
+
if isinstance(value, int) and value < 0:
|
|
67
|
+
raise ShapeMismatch(f"negative dimension {value} in shape {self.shape}", expected=">= 0", actual=value)
|
|
68
|
+
if isinstance(value, str) and not value:
|
|
69
|
+
raise ShapeMismatch(f"empty symbolic dimension in shape {self.shape}", expected="non-empty name", actual=value)
|
|
70
|
+
|
|
71
|
+
@property
|
|
72
|
+
def symbols(self) -> frozenset[str]:
|
|
73
|
+
return frozenset(value for value in self.shape if isinstance(value, str))
|
|
74
|
+
|
|
75
|
+
@property
|
|
76
|
+
def nbytes(self) -> int | None:
|
|
77
|
+
return self.nbytes_with()
|
|
78
|
+
|
|
79
|
+
def nbytes_with(self, dims: Dims | None = None) -> int | None:
|
|
80
|
+
shape = self.concrete_shape(dims)
|
|
81
|
+
if shape is None:
|
|
82
|
+
return None
|
|
83
|
+
return _product(shape) * DTYPE_BYTES[self.dtype]
|
|
84
|
+
|
|
85
|
+
def concrete_shape(self, dims: Dims | None = None) -> tuple[int, ...] | None:
|
|
86
|
+
dims = dims or {}
|
|
87
|
+
if any(isinstance(value, str) and value not in dims for value in self.shape):
|
|
88
|
+
return None
|
|
89
|
+
return tuple(dims[value] if isinstance(value, str) else value for value in self.shape)
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
@dataclass(frozen=True)
|
|
93
|
+
class ResourceContract:
|
|
94
|
+
max_memory_bytes: int | None = None
|
|
95
|
+
deterministic: bool = False
|
|
96
|
+
|
|
97
|
+
def __post_init__(self) -> None:
|
|
98
|
+
limit = self.max_memory_bytes
|
|
99
|
+
if limit is not None and (isinstance(limit, bool) or not isinstance(limit, int) or limit < 0):
|
|
100
|
+
raise ContractViolation(
|
|
101
|
+
f"max_memory_bytes must be a non-negative int, got {limit!r}",
|
|
102
|
+
expected="non-negative int", actual=limit,
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
@classmethod
|
|
106
|
+
def max_memory(cls, value: int, *, deterministic: bool = False) -> "ResourceContract":
|
|
107
|
+
return cls(value, deterministic)
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
@dataclass(frozen=True)
|
|
111
|
+
class Node:
|
|
112
|
+
name: str
|
|
113
|
+
operation: str
|
|
114
|
+
inputs: tuple[str, ...]
|
|
115
|
+
output: str
|
|
116
|
+
output_spec: TensorSpec
|
|
117
|
+
attributes: tuple[tuple[str, Any], ...] = ()
|
|
118
|
+
|
|
119
|
+
def attribute(self, name: str) -> Any:
|
|
120
|
+
return dict(self.attributes)[name]
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
@dataclass(frozen=True)
|
|
124
|
+
class MemoryValue:
|
|
125
|
+
name: str
|
|
126
|
+
bytes: int | None
|
|
127
|
+
first_use: int
|
|
128
|
+
last_use: int
|
|
129
|
+
buffer: str = ""
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
@dataclass(frozen=True)
|
|
133
|
+
class MemoryPlan:
|
|
134
|
+
values: tuple[MemoryValue, ...]
|
|
135
|
+
peak_memory_bytes: int | None
|
|
136
|
+
reusable_buffers: int
|
|
137
|
+
bookkeeping_bytes: int = 0
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
@dataclass(frozen=True)
|
|
141
|
+
class Analysis:
|
|
142
|
+
graph_name: str
|
|
143
|
+
memory_plan: MemoryPlan
|
|
144
|
+
claims: tuple[Claim, ...]
|
|
145
|
+
|
|
146
|
+
def claim(self, name: str) -> Claim:
|
|
147
|
+
return find_claim(self.claims, name)
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
@dataclass(frozen=True)
|
|
151
|
+
class ExecutionRecord:
|
|
152
|
+
graph_name: str
|
|
153
|
+
backend: str
|
|
154
|
+
runtime: str
|
|
155
|
+
deterministic: bool
|
|
156
|
+
peak_memory_bytes: int | None
|
|
157
|
+
timestamp: str
|
|
158
|
+
graph_sha256: str = ""
|
|
159
|
+
inputs_sha256: str = ""
|
|
160
|
+
output_sha256: str = ""
|
|
161
|
+
python_version: str = ""
|
|
162
|
+
claims: tuple[Claim, ...] = ()
|
|
163
|
+
|
|
164
|
+
def claim(self, name: str) -> Claim:
|
|
165
|
+
return find_claim(self.claims, name)
|
|
166
|
+
|
|
167
|
+
def to_dict(self) -> dict[str, Any]:
|
|
168
|
+
return asdict(self)
|
|
169
|
+
|
|
170
|
+
|
|
171
|
+
@dataclass(frozen=True)
|
|
172
|
+
class _Step:
|
|
173
|
+
node: Node
|
|
174
|
+
in_shapes: tuple[tuple[int, ...] | None, ...]
|
|
175
|
+
out_shape: tuple[int, ...] | None
|
|
176
|
+
out_source: str | None
|
|
177
|
+
release: tuple[str, ...]
|
|
178
|
+
|
|
179
|
+
|
|
180
|
+
@dataclass(frozen=True)
|
|
181
|
+
class _Prepared:
|
|
182
|
+
plan: MemoryPlan
|
|
183
|
+
steps: tuple[_Step, ...]
|
|
184
|
+
graph_sha256: str
|
|
185
|
+
claims: tuple[Claim, ...]
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
class Graph:
|
|
189
|
+
def __init__(self, name: str, resources: ResourceContract | None = None) -> None:
|
|
190
|
+
self.name = name
|
|
191
|
+
self.resources = resources or ResourceContract()
|
|
192
|
+
self._inputs: dict[str, TensorSpec] = {}
|
|
193
|
+
self._nodes: list[Node] = []
|
|
194
|
+
self._output: str | None = None
|
|
195
|
+
self._prepared: dict[tuple, _Prepared] = {}
|
|
196
|
+
self._lock = threading.Lock()
|
|
197
|
+
|
|
198
|
+
@property
|
|
199
|
+
def resources(self) -> ResourceContract:
|
|
200
|
+
return self._resources
|
|
201
|
+
|
|
202
|
+
@resources.setter
|
|
203
|
+
def resources(self, value: ResourceContract) -> None:
|
|
204
|
+
if not isinstance(value, ResourceContract):
|
|
205
|
+
raise ContractViolation(
|
|
206
|
+
f"{self.name}: resources must be a ResourceContract",
|
|
207
|
+
graph=self.name, expected="ResourceContract", actual=type(value).__name__,
|
|
208
|
+
)
|
|
209
|
+
self._resources = value
|
|
210
|
+
|
|
211
|
+
@property
|
|
212
|
+
def inputs(self) -> dict[str, TensorSpec]:
|
|
213
|
+
return dict(self._inputs)
|
|
214
|
+
|
|
215
|
+
@property
|
|
216
|
+
def nodes(self) -> tuple[Node, ...]:
|
|
217
|
+
return tuple(self._nodes)
|
|
218
|
+
|
|
219
|
+
def input(self, name: str, spec: TensorSpec) -> str:
|
|
220
|
+
if name in self._inputs or any(node.output == name for node in self._nodes):
|
|
221
|
+
raise ContractViolation(
|
|
222
|
+
f"{self.name}: value name already defined: {name}",
|
|
223
|
+
graph=self.name, node=name, remediation="choose a unique input name",
|
|
224
|
+
)
|
|
225
|
+
self._inputs[name] = spec
|
|
226
|
+
return name
|
|
227
|
+
|
|
228
|
+
def add(self, left: str, right: str, *, name: str = "add") -> str:
|
|
229
|
+
return self._elementwise("add", left, right, name)
|
|
230
|
+
|
|
231
|
+
def multiply(self, left: str, right: str, *, name: str = "multiply") -> str:
|
|
232
|
+
return self._elementwise("multiply", left, right, name)
|
|
233
|
+
|
|
234
|
+
def subtract(self, left: str, right: str, *, name: str = "subtract") -> str:
|
|
235
|
+
return self._elementwise("subtract", left, right, name)
|
|
236
|
+
|
|
237
|
+
def divide(self, left: str, right: str, *, name: str = "divide") -> str:
|
|
238
|
+
self._require_float(left, name)
|
|
239
|
+
return self._elementwise("divide", left, right, name)
|
|
240
|
+
|
|
241
|
+
def scale(self, value: str, factor: float, *, name: str = "scale") -> str:
|
|
242
|
+
source = self._require_float(value, name)
|
|
243
|
+
if isinstance(factor, bool) or not isinstance(factor, (int, float)):
|
|
244
|
+
raise ContractViolation(
|
|
245
|
+
f"{name}: factor must be a number", graph=self.name, node=name, expected="number", actual=type(factor).__name__,
|
|
246
|
+
)
|
|
247
|
+
output = self._unique_output(name)
|
|
248
|
+
self._nodes.append(Node(name, "scale", (value,), output, TensorSpec(source.shape, source.dtype), (("factor", float(factor)),)))
|
|
249
|
+
return output
|
|
250
|
+
|
|
251
|
+
def transpose(self, value: str, axes: tuple[int, ...] | None = None, *, name: str = "transpose") -> str:
|
|
252
|
+
source = self._spec(value)
|
|
253
|
+
rank = len(source.shape)
|
|
254
|
+
axes = tuple(reversed(range(rank))) if axes is None else tuple(axes)
|
|
255
|
+
if sorted(axes) != list(range(rank)) or any(isinstance(axis, bool) for axis in axes):
|
|
256
|
+
raise ShapeMismatch(
|
|
257
|
+
f"{name}: axes {axes} are not a permutation of {rank} dimensions",
|
|
258
|
+
graph=self.name, node=name, expected=list(range(rank)), actual=list(axes),
|
|
259
|
+
)
|
|
260
|
+
output = self._unique_output(name)
|
|
261
|
+
shape = tuple(source.shape[axis] for axis in axes)
|
|
262
|
+
self._nodes.append(Node(name, "transpose", (value,), output, TensorSpec(shape, source.dtype), (("axes", axes),)))
|
|
263
|
+
return output
|
|
264
|
+
|
|
265
|
+
def softmax(self, value: str, *, name: str = "softmax") -> str:
|
|
266
|
+
source = self._require_float(value, name)
|
|
267
|
+
self._require_rank(source, name)
|
|
268
|
+
output = self._unique_output(name)
|
|
269
|
+
self._nodes.append(Node(name, "softmax", (value,), output, TensorSpec(source.shape, source.dtype)))
|
|
270
|
+
return output
|
|
271
|
+
|
|
272
|
+
def layer_norm(self, value: str, gamma: str, beta: str, *, eps: float = 1e-5, name: str = "layer_norm") -> str:
|
|
273
|
+
source = self._require_float(value, name)
|
|
274
|
+
self._require_rank(source, name)
|
|
275
|
+
for parameter in (gamma, beta):
|
|
276
|
+
spec = self._spec(parameter)
|
|
277
|
+
if spec.shape != source.shape[-1:]:
|
|
278
|
+
raise ShapeMismatch(
|
|
279
|
+
f"{name}: {parameter} must have shape {source.shape[-1:]}",
|
|
280
|
+
graph=self.name, node=name, expected=source.shape[-1:], actual=spec.shape,
|
|
281
|
+
)
|
|
282
|
+
if spec.dtype != source.dtype:
|
|
283
|
+
raise DTypeMismatch(
|
|
284
|
+
f"{name}: {parameter} must have dtype {source.dtype}",
|
|
285
|
+
graph=self.name, node=name, expected=source.dtype, actual=spec.dtype,
|
|
286
|
+
)
|
|
287
|
+
if isinstance(eps, bool) or not isinstance(eps, (int, float)) or not eps > 0:
|
|
288
|
+
raise ContractViolation(f"{name}: eps must be positive", graph=self.name, node=name, expected="> 0", actual=eps)
|
|
289
|
+
output = self._unique_output(name)
|
|
290
|
+
self._nodes.append(Node(
|
|
291
|
+
name, "layer_norm", (value, gamma, beta), output, TensorSpec(source.shape, source.dtype), (("eps", float(eps)),),
|
|
292
|
+
))
|
|
293
|
+
return output
|
|
294
|
+
|
|
295
|
+
def relu(self, value: str, *, name: str = "relu") -> str:
|
|
296
|
+
source = self._spec(value)
|
|
297
|
+
output = self._unique_output(name)
|
|
298
|
+
self._nodes.append(Node(name, "relu", (value,), output, TensorSpec(source.shape, source.dtype, source.layout)))
|
|
299
|
+
return output
|
|
300
|
+
|
|
301
|
+
def reshape(self, value: str, shape: tuple[Any, ...], *, name: str = "reshape") -> str:
|
|
302
|
+
source = self._spec(value)
|
|
303
|
+
target = TensorSpec(shape, source.dtype, source.layout, source.role)
|
|
304
|
+
if _element_count(source.shape) != _element_count(target.shape):
|
|
305
|
+
raise ShapeMismatch(
|
|
306
|
+
f"{name}: cannot prove reshape {source.shape} -> {target.shape} keeps element count",
|
|
307
|
+
graph=self.name, node=name, expected=source.shape, actual=target.shape,
|
|
308
|
+
)
|
|
309
|
+
output = self._unique_output(name)
|
|
310
|
+
self._nodes.append(Node(name, "reshape", (value,), output, target))
|
|
311
|
+
return output
|
|
312
|
+
|
|
313
|
+
def matmul(self, left: str, right: str, *, name: str = "matmul") -> str:
|
|
314
|
+
left_spec, right_spec = self._spec(left), self._spec(right)
|
|
315
|
+
if len(left_spec.shape) != 2 or len(right_spec.shape) != 2:
|
|
316
|
+
raise ShapeMismatch(
|
|
317
|
+
f"{name}: matmul requires rank-2 tensors",
|
|
318
|
+
graph=self.name, node=name, expected="rank 2", actual=(left_spec.shape, right_spec.shape),
|
|
319
|
+
)
|
|
320
|
+
if left_spec.shape[1] != right_spec.shape[0]:
|
|
321
|
+
raise ShapeMismatch(
|
|
322
|
+
f"{name}: incompatible inner dimensions",
|
|
323
|
+
graph=self.name, node=name, expected=left_spec.shape[1], actual=right_spec.shape[0],
|
|
324
|
+
)
|
|
325
|
+
if left_spec.dtype != right_spec.dtype:
|
|
326
|
+
raise DTypeMismatch(
|
|
327
|
+
f"{name}: matmul requires matching dtypes",
|
|
328
|
+
graph=self.name, node=name, expected=left_spec.dtype, actual=right_spec.dtype,
|
|
329
|
+
)
|
|
330
|
+
output = self._unique_output(name)
|
|
331
|
+
self._nodes.append(Node(name, "matmul", (left, right), output, TensorSpec((left_spec.shape[0], right_spec.shape[1]), left_spec.dtype)))
|
|
332
|
+
return output
|
|
333
|
+
|
|
334
|
+
def output(self, value: str) -> None:
|
|
335
|
+
self._spec(value)
|
|
336
|
+
self._output = value
|
|
337
|
+
|
|
338
|
+
def validate(self, dims: Dims | None = None) -> None:
|
|
339
|
+
if self._output is None:
|
|
340
|
+
raise UnknownValue(f"graph {self.name} has no output", graph=self.name, remediation="call graph.output(value)")
|
|
341
|
+
if self.resources.max_memory_bytes is not None:
|
|
342
|
+
self._check_memory(self.memory_plan(dims), dims)
|
|
343
|
+
|
|
344
|
+
def _check_memory(self, plan: MemoryPlan, dims: Dims | None) -> None:
|
|
345
|
+
limit = self.resources.max_memory_bytes
|
|
346
|
+
if limit is None:
|
|
347
|
+
return
|
|
348
|
+
if plan.peak_memory_bytes is None:
|
|
349
|
+
unbound = sorted(self._symbols() - set(dims or {}))
|
|
350
|
+
raise ContractViolation(
|
|
351
|
+
f"{self.name}: cannot verify memory limit {limit} with unbound symbolic dimensions {unbound}",
|
|
352
|
+
graph=self.name, expected="bound dimensions", actual=unbound,
|
|
353
|
+
remediation="pass dims to validate() or execute with concrete inputs",
|
|
354
|
+
)
|
|
355
|
+
if plan.peak_memory_bytes > limit:
|
|
356
|
+
raise ResourceLimitExceeded(
|
|
357
|
+
f"{self.name}: peak memory {plan.peak_memory_bytes} exceeds limit {limit}",
|
|
358
|
+
graph=self.name, expected=limit, actual=plan.peak_memory_bytes,
|
|
359
|
+
remediation="raise max_memory_bytes or reduce tensor sizes",
|
|
360
|
+
)
|
|
361
|
+
|
|
362
|
+
def analyze(self, dims: Dims | None = None) -> Analysis:
|
|
363
|
+
plan = self.memory_plan(dims)
|
|
364
|
+
output = self._spec(self._output) if self._output is not None else None
|
|
365
|
+
return Analysis(self.name, plan, (
|
|
366
|
+
Claim("shapes_consistent", True, Basis.PROVEN),
|
|
367
|
+
Claim("output_spec", output, Basis.PROVEN),
|
|
368
|
+
Claim("dimensions", dict(dims or {}), Basis.DECLARED),
|
|
369
|
+
Claim("unbound_dimensions", sorted(self._symbols() - set(dims or {})), Basis.PROVEN),
|
|
370
|
+
*self._contract_claims(plan),
|
|
371
|
+
))
|
|
372
|
+
|
|
373
|
+
def memory_plan(self, dims: Dims | None = None) -> MemoryPlan:
|
|
374
|
+
steps = len(self._nodes)
|
|
375
|
+
first_use = {name: 0 for name in self._inputs}
|
|
376
|
+
last_use = {name: steps for name in self._inputs}
|
|
377
|
+
for index, node in enumerate(self._nodes, start=1):
|
|
378
|
+
first_use[node.output] = index
|
|
379
|
+
last_use[node.output] = steps if node.output == self._output else index
|
|
380
|
+
for name in node.inputs:
|
|
381
|
+
last_use[name] = max(last_use[name], index)
|
|
382
|
+
|
|
383
|
+
buffer = {name: name for name in self._inputs}
|
|
384
|
+
members: dict[str, list[str]] = {name: [name] for name in self._inputs}
|
|
385
|
+
contiguous = {name: True for name in self._inputs}
|
|
386
|
+
owned: set[str] = set()
|
|
387
|
+
scratch = [0] * (steps + 1)
|
|
388
|
+
for index, node in enumerate(self._nodes, start=1):
|
|
389
|
+
target = None
|
|
390
|
+
source = node.inputs[0]
|
|
391
|
+
if node.operation == "transpose":
|
|
392
|
+
target = buffer[source]
|
|
393
|
+
contiguous[node.output] = node.attribute("axes") == tuple(range(len(node.output_spec.shape)))
|
|
394
|
+
elif node.operation == "reshape" and contiguous[source]:
|
|
395
|
+
target = buffer[source]
|
|
396
|
+
contiguous[node.output] = True
|
|
397
|
+
else:
|
|
398
|
+
contiguous[node.output] = True
|
|
399
|
+
if node.operation in IN_PLACE:
|
|
400
|
+
for name in node.inputs:
|
|
401
|
+
candidate = buffer[name]
|
|
402
|
+
if (candidate in owned and contiguous[name]
|
|
403
|
+
and self._spec(name).shape == node.output_spec.shape
|
|
404
|
+
and max(last_use[member] for member in members[candidate]) == index):
|
|
405
|
+
target = candidate
|
|
406
|
+
break
|
|
407
|
+
if target is None:
|
|
408
|
+
target = node.output
|
|
409
|
+
members[target] = []
|
|
410
|
+
owned.add(target)
|
|
411
|
+
buffer[node.output] = target
|
|
412
|
+
members[target].append(node.output)
|
|
413
|
+
scratch[index] = self._scratch_bytes(node, dims)
|
|
414
|
+
|
|
415
|
+
reports = tuple(
|
|
416
|
+
MemoryValue(name, self._spec(name).nbytes_with(dims), first_use[name], last_use[name], buffer[name])
|
|
417
|
+
for name in first_use
|
|
418
|
+
)
|
|
419
|
+
sizes = {name: self._spec(name).nbytes_with(dims) for name in members}
|
|
420
|
+
if any(size is None for size in sizes.values()) or None in scratch:
|
|
421
|
+
peak = None
|
|
422
|
+
else:
|
|
423
|
+
spans = {
|
|
424
|
+
name: (first_use[name], max(last_use[member] for member in group))
|
|
425
|
+
for name, group in members.items()
|
|
426
|
+
}
|
|
427
|
+
peak = max(
|
|
428
|
+
(sum(sizes[name] for name, (first, last) in spans.items() if first <= step <= last) + scratch[step]
|
|
429
|
+
for step in range(steps + 1)),
|
|
430
|
+
default=0,
|
|
431
|
+
)
|
|
432
|
+
reusable = sum(value.buffer != value.name for value in reports)
|
|
433
|
+
bookkeeping = (
|
|
434
|
+
BOOKKEEPING_BASE_BYTES
|
|
435
|
+
+ BOOKKEEPING_INPUT_BYTES * len(self._inputs)
|
|
436
|
+
+ BOOKKEEPING_NODE_BYTES * len(self._nodes)
|
|
437
|
+
)
|
|
438
|
+
return MemoryPlan(reports, None if peak is None else peak + bookkeeping, reusable, bookkeeping)
|
|
439
|
+
|
|
440
|
+
def execute(
|
|
441
|
+
self, values: Mapping[str, Any], *, backend: str = "python", digests: bool = False,
|
|
442
|
+
) -> tuple[Any, ExecutionRecord]:
|
|
443
|
+
from .backends import get_backend
|
|
444
|
+
|
|
445
|
+
impl = get_backend(backend)
|
|
446
|
+
try:
|
|
447
|
+
return self._execute(impl, values, digests)
|
|
448
|
+
except GraphspaceError as error:
|
|
449
|
+
if error.graph is None:
|
|
450
|
+
error.graph = self.name
|
|
451
|
+
raise
|
|
452
|
+
|
|
453
|
+
def _execute(self, impl: Any, values: Mapping[str, Any], digests: bool) -> tuple[Any, ExecutionRecord]:
|
|
454
|
+
missing = sorted(set(self._inputs) - set(values))
|
|
455
|
+
unexpected = sorted(set(values) - set(self._inputs))
|
|
456
|
+
if missing:
|
|
457
|
+
raise UnknownValue(
|
|
458
|
+
f"{self.name}: missing input values: {missing}",
|
|
459
|
+
expected=sorted(self._inputs), actual=sorted(values), remediation=f"provide values for {missing}",
|
|
460
|
+
)
|
|
461
|
+
if unexpected:
|
|
462
|
+
raise UnknownValue(
|
|
463
|
+
f"{self.name}: values for undeclared inputs: {unexpected}",
|
|
464
|
+
expected=sorted(self._inputs), actual=sorted(values), remediation=f"remove {unexpected}",
|
|
465
|
+
)
|
|
466
|
+
dims = self._infer_dims({name: impl.length(name, data) for name, data in values.items()})
|
|
467
|
+
if self._output is None:
|
|
468
|
+
self.validate(dims)
|
|
469
|
+
prepared = self._prepare(dims)
|
|
470
|
+
self._check_memory(prepared.plan, dims)
|
|
471
|
+
computed = {name: impl.coerce(name, values[name], spec, dims) for name, spec in self._inputs.items()}
|
|
472
|
+
digest_claims = []
|
|
473
|
+
if digests:
|
|
474
|
+
digest_claims.append(Claim("inputs_sha256", impl.digest(
|
|
475
|
+
{name: (spec.dtype, computed[name]) for name, spec in self._inputs.items()}
|
|
476
|
+
), Basis.MEASURED))
|
|
477
|
+
with impl.session():
|
|
478
|
+
for step in prepared.steps:
|
|
479
|
+
node = step.node
|
|
480
|
+
out = computed[step.out_source] if step.out_source is not None else None
|
|
481
|
+
computed[node.output] = impl.run(node, [computed[name] for name in node.inputs], step.in_shapes, step.out_shape, out)
|
|
482
|
+
for name in step.release:
|
|
483
|
+
del computed[name]
|
|
484
|
+
result = computed[self._output]
|
|
485
|
+
if digests:
|
|
486
|
+
digest_claims.append(Claim("output_sha256", impl.digest(
|
|
487
|
+
{self._output: (self._spec(self._output).dtype, result)}
|
|
488
|
+
), Basis.MEASURED))
|
|
489
|
+
return result, ExecutionRecord(
|
|
490
|
+
self.name, impl.name, f"graphspace-{__version__}", self.resources.deterministic,
|
|
491
|
+
prepared.plan.peak_memory_bytes, datetime.now(timezone.utc).isoformat(),
|
|
492
|
+
graph_sha256=prepared.graph_sha256,
|
|
493
|
+
inputs_sha256=digest_claims[0].value if digests else "",
|
|
494
|
+
output_sha256=digest_claims[1].value if digests else "",
|
|
495
|
+
python_version=PYTHON_VERSION,
|
|
496
|
+
claims=(
|
|
497
|
+
Claim("dimensions", dict(dims), Basis.INFERRED),
|
|
498
|
+
*prepared.claims,
|
|
499
|
+
*digest_claims,
|
|
500
|
+
Claim("backend_version", impl.version, Basis.BACKEND_REPORTED),
|
|
501
|
+
),
|
|
502
|
+
)
|
|
503
|
+
|
|
504
|
+
def _prepare(self, dims: Dims) -> _Prepared:
|
|
505
|
+
key = (len(self._inputs), len(self._nodes), self._output, self.resources, tuple(sorted(dims.items())))
|
|
506
|
+
with self._lock:
|
|
507
|
+
if key in self._prepared:
|
|
508
|
+
return self._prepared[key]
|
|
509
|
+
plan = self.memory_plan(dims)
|
|
510
|
+
buffer = {value.name: value.buffer for value in plan.values}
|
|
511
|
+
last_use = {value.name: value.last_use for value in plan.values}
|
|
512
|
+
steps = []
|
|
513
|
+
for index, node in enumerate(self._nodes, start=1):
|
|
514
|
+
out_source = None
|
|
515
|
+
if node.operation not in ("reshape", "transpose") and buffer[node.output] != node.output:
|
|
516
|
+
out_source = next(name for name in node.inputs if buffer[name] == buffer[node.output])
|
|
517
|
+
release = tuple(
|
|
518
|
+
name for name in dict.fromkeys((*node.inputs, node.output))
|
|
519
|
+
if last_use[name] == index and name != self._output and name not in self._inputs
|
|
520
|
+
)
|
|
521
|
+
steps.append(_Step(
|
|
522
|
+
node,
|
|
523
|
+
tuple(self._spec(name).concrete_shape(dims) for name in node.inputs),
|
|
524
|
+
node.output_spec.concrete_shape(dims),
|
|
525
|
+
out_source,
|
|
526
|
+
release,
|
|
527
|
+
))
|
|
528
|
+
dtypes = {spec.dtype for spec in [*self._inputs.values(), *(node.output_spec for node in self._nodes)]}
|
|
529
|
+
claims = (
|
|
530
|
+
Claim("inputs_valid", True, Basis.RUNTIME_CHECKED),
|
|
531
|
+
*([Claim("integer_range", True, Basis.RUNTIME_CHECKED)] if dtypes & set(INT_RANGES) else []),
|
|
532
|
+
*self._contract_claims(plan),
|
|
533
|
+
)
|
|
534
|
+
prepared = _Prepared(plan, tuple(steps), _sha256(self._describe()), claims)
|
|
535
|
+
with self._lock:
|
|
536
|
+
if key not in self._prepared and len(self._prepared) >= PREPARED_CACHE_SIZE:
|
|
537
|
+
del self._prepared[next(iter(self._prepared))]
|
|
538
|
+
self._prepared[key] = prepared
|
|
539
|
+
return prepared
|
|
540
|
+
|
|
541
|
+
def _contract_claims(self, plan: MemoryPlan) -> list[Claim]:
|
|
542
|
+
limit = self.resources.max_memory_bytes
|
|
543
|
+
peak = plan.peak_memory_bytes
|
|
544
|
+
within = None if limit is None or peak is None else peak <= limit
|
|
545
|
+
return [
|
|
546
|
+
Claim("peak_memory_bytes", peak, Basis.ESTIMATED),
|
|
547
|
+
Claim("bookkeeping_bytes", plan.bookkeeping_bytes, Basis.ESTIMATED),
|
|
548
|
+
Claim("max_memory_bytes", limit, Basis.DECLARED),
|
|
549
|
+
Claim("within_memory_limit", within, Basis.ESTIMATED),
|
|
550
|
+
Claim("deterministic", self.resources.deterministic, Basis.DECLARED),
|
|
551
|
+
]
|
|
552
|
+
|
|
553
|
+
def _elementwise(self, operation: str, left: str, right: str, name: str) -> str:
|
|
554
|
+
spec_left, spec_right = self._spec(left), self._spec(right)
|
|
555
|
+
shape = _broadcast(spec_left.shape, spec_right.shape)
|
|
556
|
+
if shape is None:
|
|
557
|
+
raise ShapeMismatch(
|
|
558
|
+
f"{name}: cannot broadcast {spec_left.shape} with {spec_right.shape}",
|
|
559
|
+
graph=self.name, node=name, expected=spec_left.shape, actual=spec_right.shape,
|
|
560
|
+
)
|
|
561
|
+
if spec_left.dtype != spec_right.dtype:
|
|
562
|
+
raise DTypeMismatch(
|
|
563
|
+
f"{name}: {operation} requires identical dtypes: {spec_left.dtype} != {spec_right.dtype}",
|
|
564
|
+
graph=self.name, node=name, expected=spec_left.dtype, actual=spec_right.dtype,
|
|
565
|
+
)
|
|
566
|
+
if spec_left.layout != spec_right.layout:
|
|
567
|
+
raise ShapeMismatch(
|
|
568
|
+
f"{name}: {operation} requires identical layouts: {spec_left.layout} != {spec_right.layout}",
|
|
569
|
+
graph=self.name, node=name, expected=spec_left.layout, actual=spec_right.layout,
|
|
570
|
+
)
|
|
571
|
+
output = self._unique_output(name)
|
|
572
|
+
self._nodes.append(Node(name, operation, (left, right), output, TensorSpec(shape, spec_left.dtype, spec_left.layout)))
|
|
573
|
+
return output
|
|
574
|
+
|
|
575
|
+
def _scratch_bytes(self, node: Node, dims: Dims | None) -> int | None:
|
|
576
|
+
shape = node.output_spec.concrete_shape(dims)
|
|
577
|
+
if shape is None:
|
|
578
|
+
return None
|
|
579
|
+
if node.operation in UNBUFFERED:
|
|
580
|
+
return 0
|
|
581
|
+
itemsize = DTYPE_BYTES[node.output_spec.dtype]
|
|
582
|
+
rows = 2 * _product(shape[:-1]) * itemsize if node.operation in ROW_SCRATCH else 0
|
|
583
|
+
return rows + (len(node.inputs) + 1) * min(_product(shape), UFUNC_BUFFER_ELEMENTS) * itemsize
|
|
584
|
+
|
|
585
|
+
def _require_float(self, value: str, name: str) -> TensorSpec:
|
|
586
|
+
spec = self._spec(value)
|
|
587
|
+
if spec.dtype not in FLOAT_DTYPES:
|
|
588
|
+
raise DTypeMismatch(
|
|
589
|
+
f"{name}: requires a float dtype, got {spec.dtype}",
|
|
590
|
+
graph=self.name, node=name, expected=sorted(FLOAT_DTYPES), actual=spec.dtype,
|
|
591
|
+
)
|
|
592
|
+
return spec
|
|
593
|
+
|
|
594
|
+
def _require_rank(self, spec: TensorSpec, name: str) -> None:
|
|
595
|
+
if not spec.shape:
|
|
596
|
+
raise ShapeMismatch(f"{name}: requires rank 1 or higher", graph=self.name, node=name, expected=">= 1", actual=0)
|
|
597
|
+
|
|
598
|
+
def _infer_dims(self, lengths: Mapping[str, int]) -> dict[str, int]:
|
|
599
|
+
dims: dict[str, int] = {}
|
|
600
|
+
pending = [name for name, spec in self._inputs.items() if spec.symbols]
|
|
601
|
+
while pending:
|
|
602
|
+
progressed = False
|
|
603
|
+
for name in list(pending):
|
|
604
|
+
spec = self._inputs[name]
|
|
605
|
+
unbound = [value for value in spec.shape if isinstance(value, str) and value not in dims]
|
|
606
|
+
if not unbound:
|
|
607
|
+
pending.remove(name)
|
|
608
|
+
continue
|
|
609
|
+
if len(unbound) > 1:
|
|
610
|
+
continue
|
|
611
|
+
known = _product(dims.get(value, value) for value in spec.shape if value != unbound[0])
|
|
612
|
+
length = lengths[name]
|
|
613
|
+
if known == 0 or length % known:
|
|
614
|
+
raise ShapeMismatch(
|
|
615
|
+
f"{self.name}: input {name} has {length} elements, incompatible with shape {spec.shape}",
|
|
616
|
+
node=name, expected=spec.shape, actual=length,
|
|
617
|
+
)
|
|
618
|
+
dims[unbound[0]] = length // known
|
|
619
|
+
pending.remove(name)
|
|
620
|
+
progressed = True
|
|
621
|
+
if pending and not progressed:
|
|
622
|
+
raise ShapeMismatch(
|
|
623
|
+
f"{self.name}: cannot infer symbolic dimensions for inputs {sorted(pending)}",
|
|
624
|
+
expected="one unbound dimension per input", actual=sorted(pending),
|
|
625
|
+
)
|
|
626
|
+
return dims
|
|
627
|
+
|
|
628
|
+
def _symbols(self) -> set[str]:
|
|
629
|
+
symbols: set[str] = set()
|
|
630
|
+
for spec in [*self._inputs.values(), *(node.output_spec for node in self._nodes)]:
|
|
631
|
+
symbols |= spec.symbols
|
|
632
|
+
return symbols
|
|
633
|
+
|
|
634
|
+
def _describe(self) -> dict[str, Any]:
|
|
635
|
+
return {
|
|
636
|
+
"name": self.name,
|
|
637
|
+
"resources": [self.resources.max_memory_bytes, self.resources.deterministic],
|
|
638
|
+
"inputs": [[name, list(spec.shape), spec.dtype, spec.layout, spec.role] for name, spec in self._inputs.items()],
|
|
639
|
+
"nodes": [
|
|
640
|
+
[node.name, node.operation, list(node.inputs), node.output, list(node.output_spec.shape),
|
|
641
|
+
node.output_spec.dtype, [[key, list(value) if isinstance(value, tuple) else value] for key, value in node.attributes]]
|
|
642
|
+
for node in self._nodes
|
|
643
|
+
],
|
|
644
|
+
"output": self._output,
|
|
645
|
+
}
|
|
646
|
+
|
|
647
|
+
def _spec(self, value: str) -> TensorSpec:
|
|
648
|
+
if value in self._inputs:
|
|
649
|
+
return self._inputs[value]
|
|
650
|
+
for node in reversed(self._nodes):
|
|
651
|
+
if node.output == value:
|
|
652
|
+
return node.output_spec
|
|
653
|
+
raise UnknownValue(f"unknown graph value: {value}", graph=self.name, node=value)
|
|
654
|
+
|
|
655
|
+
def _unique_output(self, name: str) -> str:
|
|
656
|
+
existing = set(self._inputs) | {node.output for node in self._nodes}
|
|
657
|
+
output = name
|
|
658
|
+
suffix = 1
|
|
659
|
+
while output in existing:
|
|
660
|
+
suffix += 1
|
|
661
|
+
output = f"{name}_{suffix}"
|
|
662
|
+
return output
|
|
663
|
+
|
|
664
|
+
|
|
665
|
+
def _product(values) -> int:
|
|
666
|
+
result = 1
|
|
667
|
+
for value in values:
|
|
668
|
+
result *= value
|
|
669
|
+
return result
|
|
670
|
+
|
|
671
|
+
|
|
672
|
+
def _broadcast(left: tuple[Any, ...], right: tuple[Any, ...]) -> tuple[Any, ...] | None:
|
|
673
|
+
rank = max(len(left), len(right))
|
|
674
|
+
left = (1,) * (rank - len(left)) + tuple(left)
|
|
675
|
+
right = (1,) * (rank - len(right)) + tuple(right)
|
|
676
|
+
shape = []
|
|
677
|
+
for a, b in zip(left, right):
|
|
678
|
+
if a == b or b == 1:
|
|
679
|
+
shape.append(a)
|
|
680
|
+
elif a == 1:
|
|
681
|
+
shape.append(b)
|
|
682
|
+
else:
|
|
683
|
+
return None
|
|
684
|
+
return tuple(shape)
|
|
685
|
+
|
|
686
|
+
|
|
687
|
+
def _element_count(shape: tuple[Any, ...]) -> tuple[int, tuple[str, ...]]:
|
|
688
|
+
return (
|
|
689
|
+
_product(value for value in shape if isinstance(value, int)),
|
|
690
|
+
tuple(sorted(value for value in shape if isinstance(value, str))),
|
|
691
|
+
)
|
|
692
|
+
|
|
693
|
+
|
|
694
|
+
def _sha256(payload: Any) -> str:
|
|
695
|
+
encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()
|
|
696
|
+
return hashlib.sha256(encoded).hexdigest()
|
|
697
|
+
|