-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy patharray_interop.py
More file actions
50 lines (43 loc) · 1.95 KB
/
Copy patharray_interop.py
File metadata and controls
50 lines (43 loc) · 1.95 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
"""Purpose: one operation set for every DLPack library: NumPy flows through
the same MeTTa functions torch does, a mixed call converts through DLPack,
and DLTensor is a protocol type the engine really checks.
Open Obligations:
To Do: None
Hacks: None
Future Enhancements: None
"""
from _common import check, done, skip
try:
# The array layer is a package of its own, metta-arrays, and which library
# is its default is a row metta-numpy registers; both arrive with
# `pip install 'pymetta[arrays]'`.
import array_api_compat # noqa: F401
import metta_arrays as arrays
import metta_numpy # noqa: F401 -- the row that makes numpy the default
import numpy
except ImportError:
skip("metta-arrays, numpy and array-api-compat are needed")
from metta import Expression, MeTTa, S, V, convert, ground
m = MeTTa().space()
arrays.install(m, default=numpy)
check("matmul over numpy",
m.run("!(t-tolist (matmul (tensor ((1.0 2.0))) (tensor ((3.0) (4.0)))))"),
[[Expression(Expression(11.0))]])
# The tensor is built first and the type read off the VALUE. get-type does
# not evaluate its argument, so asking about the unreduced call `(tensor
# (1.0))` reports what the expression is declared to be, not what building it
# would produce; both arbiters answer %Undefined% there.
(types,) = m.run("!(collapse (let $t (tensor (1.0)) (get-type $t)))")
check("protocol typing", S.DLTensor in list(types[0]))
array = numpy.arange(4.0)
m.add(S.holds(ground(array)))
check("identity through the space", convert.decode(m.match(S.holds(V.a))[0].a) is array)
try:
import torch
left, right = numpy.ones((2, 2), dtype=numpy.float32), torch.ones(2, 2)
m.add(S.pair(ground(left), ground(right)))
(out,) = m.run("!(t-item (t-sum (match (context-space) (pair $a $b) (matmul $a $b))))")
check("mixed numpy@torch via DLPack", float(out[0]), 8.0)
except ImportError:
print(" (torch absent: mixed-library half skipped)")
done("array_interop")