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/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
+