metatensor-core 0.2.2__tar.gz → 0.2.4__tar.gz
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.
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/CMakeLists.txt +1 -1
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/PKG-INFO +2 -1
- metatensor_core-0.2.4/git_version_info +2 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_block.py +90 -43
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_c_api.py +22 -129
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_data/_array.py +17 -109
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_data/_extract.py +6 -10
- metatensor_core-0.2.4/metatensor/_html.py +425 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_labels.py +126 -84
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_status.py +1 -1
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_tensor.py +170 -10
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/io/_block.py +1 -2
- metatensor_core-0.2.4/metatensor-core-cxx-0.2.4.tar.gz +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor_core.egg-info/PKG-INFO +2 -1
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor_core.egg-info/SOURCES.txt +2 -2
- metatensor_core-0.2.4/metatensor_core.egg-info/requires.txt +2 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/pyproject.toml +1 -1
- metatensor_core-0.2.2/git_version_info +0 -2
- metatensor_core-0.2.2/metatensor/_data/_dlpack.py +0 -293
- metatensor_core-0.2.2/metatensor-core-cxx-0.2.2.tar.gz +0 -0
- metatensor_core-0.2.2/metatensor_core.egg-info/requires.txt +0 -1
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/AUTHORS +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/LICENSE +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/MANIFEST.in +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/README.rst +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/__init__.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_c_lib.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_data/__init__.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_version.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/io/__init__.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/io/_labels.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/io/_tensor.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/io/_utils.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/learn.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/operations.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/torch.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/utils.py +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor_core.egg-info/dependency_links.txt +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor_core.egg-info/not-zip-safe +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor_core.egg-info/top_level.txt +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/setup.cfg +0 -0
- {metatensor_core-0.2.2 → metatensor_core-0.2.4}/setup.py +0 -0
|
@@ -14,7 +14,7 @@ set(METATENSOR_CORE_SOURCE_DIR "" CACHE PATH "Path to the sources of metatensor-
|
|
|
14
14
|
|
|
15
15
|
file(REMOVE ${CMAKE_INSTALL_PREFIX}/_external.py)
|
|
16
16
|
|
|
17
|
-
set(REQUIRED_METATENSOR_VERSION "0.2.
|
|
17
|
+
set(REQUIRED_METATENSOR_VERSION "0.2.4")
|
|
18
18
|
if(${METATENSOR_CORE_PYTHON_USE_EXTERNAL_LIB})
|
|
19
19
|
# when building a source checkout, update version to include git information
|
|
20
20
|
# this will not apply when building a sdist
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: metatensor-core
|
|
3
|
-
Version: 0.2.
|
|
3
|
+
Version: 0.2.4
|
|
4
4
|
Summary: Python bindings for metatensor
|
|
5
5
|
Author: Guillaume Fraux, Rohit Goswami, Philip Loche, Filippo Bigi, Joseph W. Abbott, Davide Tisi, Alexander Goscinski, Michele Ceriotti
|
|
6
6
|
License-Expression: BSD-3-Clause
|
|
@@ -27,6 +27,7 @@ Description-Content-Type: text/x-rst
|
|
|
27
27
|
License-File: LICENSE
|
|
28
28
|
License-File: AUTHORS
|
|
29
29
|
Requires-Dist: numpy
|
|
30
|
+
Requires-Dist: ctypes-dlpack
|
|
30
31
|
Dynamic: author
|
|
31
32
|
Dynamic: license-file
|
|
32
33
|
|
|
@@ -48,17 +48,14 @@ class TensorBlock:
|
|
|
48
48
|
... properties=Labels("properties", np.array([[0], [1], [2]])),
|
|
49
49
|
... )
|
|
50
50
|
>>> block
|
|
51
|
-
TensorBlock
|
|
52
|
-
samples
|
|
53
|
-
|
|
54
|
-
properties (3): ['properties']
|
|
55
|
-
gradients: None
|
|
51
|
+
TensorBlock with shape (2, 3)
|
|
52
|
+
samples: [samples]
|
|
53
|
+
properties: [properties]
|
|
56
54
|
>>> block.samples
|
|
57
|
-
Labels
|
|
55
|
+
Labels
|
|
58
56
|
samples
|
|
59
57
|
4
|
|
60
58
|
2
|
|
61
|
-
)
|
|
62
59
|
>>> block.values[block.samples.position([2])]
|
|
63
60
|
array([3, 5, 6])
|
|
64
61
|
"""
|
|
@@ -254,31 +251,82 @@ class TensorBlock:
|
|
|
254
251
|
# The block has been released
|
|
255
252
|
return "TensorBlock(<empty>)"
|
|
256
253
|
|
|
254
|
+
def _format_names(names):
|
|
255
|
+
return "[" + ", ".join(names) + "]"
|
|
256
|
+
|
|
257
257
|
if len(self._gradient_parameters) != 0:
|
|
258
|
-
s =
|
|
259
|
-
|
|
260
|
-
|
|
261
|
-
|
|
262
|
-
|
|
263
|
-
s += " components ("
|
|
264
|
-
s += ", ".join([str(len(c)) for c in self.components])
|
|
265
|
-
s += "): ["
|
|
266
|
-
for ic in self.components:
|
|
267
|
-
for name in ic.names[:]:
|
|
268
|
-
s += "'" + name + "', "
|
|
269
|
-
if len(self.components) > 0:
|
|
270
|
-
s = s[:-2]
|
|
271
|
-
s += "]\n"
|
|
272
|
-
s += f" properties ({len(self.properties)}): "
|
|
273
|
-
s += f"{str(list(self.properties.names))}\n"
|
|
274
|
-
s += " gradients: "
|
|
275
|
-
if len(self.gradients_list()) > 0:
|
|
276
|
-
s += f"{str(list(self.gradients_list()))}"
|
|
258
|
+
s = (
|
|
259
|
+
f"TensorBlock gradient for "
|
|
260
|
+
f"'{'/'.join(self._gradient_parameters)}', "
|
|
261
|
+
f"with shape {tuple(self.values.shape)}\n"
|
|
262
|
+
)
|
|
277
263
|
else:
|
|
278
|
-
s
|
|
264
|
+
s = f"TensorBlock with shape {tuple(self.values.shape)}\n"
|
|
265
|
+
|
|
266
|
+
s += f" samples: {_format_names(self.samples.names)}\n"
|
|
267
|
+
|
|
268
|
+
component_names = []
|
|
269
|
+
for component in self.components:
|
|
270
|
+
component_names.extend(component.names)
|
|
271
|
+
if len(component_names) != 0:
|
|
272
|
+
s += f" components: {_format_names(component_names)}\n"
|
|
273
|
+
|
|
274
|
+
s += f" properties: {_format_names(self.properties.names)}"
|
|
275
|
+
|
|
276
|
+
gradients = self.gradients_list()
|
|
277
|
+
if len(gradients) != 0:
|
|
278
|
+
s += "\n\n gradients:"
|
|
279
|
+
max_len = max(len(p) for p in gradients)
|
|
280
|
+
for parameter in gradients:
|
|
281
|
+
gradient = self.gradient(parameter)
|
|
282
|
+
grad_shape = tuple(int(v) for v in gradient.values.shape)
|
|
283
|
+
s += (
|
|
284
|
+
f"\n {parameter.ljust(max_len)} "
|
|
285
|
+
f"=> TensorBlock with shape {grad_shape}"
|
|
286
|
+
)
|
|
279
287
|
|
|
280
288
|
return s
|
|
281
289
|
|
|
290
|
+
def _repr_html_(self) -> str:
|
|
291
|
+
"""HTML representation for Jupyter notebooks."""
|
|
292
|
+
if not self._ptr:
|
|
293
|
+
# The block has been released
|
|
294
|
+
return "TensorBlock(<empty>)"
|
|
295
|
+
|
|
296
|
+
from metatensor._html import TensorBlockData, block_html
|
|
297
|
+
|
|
298
|
+
def _block_data(block):
|
|
299
|
+
gradients = {}
|
|
300
|
+
for parameter in block.gradients_list():
|
|
301
|
+
grad = block.gradient(parameter)
|
|
302
|
+
gradients[parameter] = _block_data(grad)
|
|
303
|
+
|
|
304
|
+
return TensorBlockData(
|
|
305
|
+
values_shape=tuple(block.values.shape),
|
|
306
|
+
samples=(block.samples.names, block.samples.values),
|
|
307
|
+
components=[(c.names, c.values) for c in block.components],
|
|
308
|
+
properties=(block.properties.names, block.properties.values),
|
|
309
|
+
gradients=gradients,
|
|
310
|
+
)
|
|
311
|
+
|
|
312
|
+
data = _block_data(self)
|
|
313
|
+
|
|
314
|
+
body = block_html(data, module="metatensor")
|
|
315
|
+
|
|
316
|
+
class_name = "metatensor.TensorBlock"
|
|
317
|
+
if len(self._gradient_parameters) != 0:
|
|
318
|
+
gradient_for = "/".join(self._gradient_parameters)
|
|
319
|
+
rest = (
|
|
320
|
+
f" gradient for '{gradient_for}', with shape {tuple(self.values.shape)}"
|
|
321
|
+
)
|
|
322
|
+
else:
|
|
323
|
+
rest = f" with shape {tuple(self.values.shape)}"
|
|
324
|
+
|
|
325
|
+
header = f"<div><strong>{class_name}</strong>{rest}</div>"
|
|
326
|
+
hr = "<hr style='border: none; border-top: 1px solid #888; margin: 0;'>"
|
|
327
|
+
|
|
328
|
+
return header + hr + body
|
|
329
|
+
|
|
282
330
|
def __eq__(self, other):
|
|
283
331
|
from metatensor.operations import equal_block
|
|
284
332
|
|
|
@@ -398,19 +446,16 @@ class TensorBlock:
|
|
|
398
446
|
|
|
399
447
|
>>> positions_gradient = block.gradient("positions")
|
|
400
448
|
>>> print(positions_gradient)
|
|
401
|
-
|
|
402
|
-
samples
|
|
403
|
-
components
|
|
404
|
-
properties
|
|
405
|
-
gradients: None
|
|
406
|
-
|
|
449
|
+
TensorBlock gradient for 'positions', with shape (2, 3, 1, 5)
|
|
450
|
+
samples: [sample, atom]
|
|
451
|
+
components: [direction, component]
|
|
452
|
+
properties: [property]
|
|
407
453
|
>>> cell_gradient = block.gradient("cell")
|
|
408
454
|
>>> print(cell_gradient)
|
|
409
|
-
|
|
410
|
-
samples
|
|
411
|
-
components
|
|
412
|
-
properties
|
|
413
|
-
gradients: None
|
|
455
|
+
TensorBlock gradient for 'cell', with shape (2, 3, 3, 1, 5)
|
|
456
|
+
samples: [sample]
|
|
457
|
+
components: [direction_1, direction_2, component]
|
|
458
|
+
properties: [property]
|
|
414
459
|
"""
|
|
415
460
|
gradient_block = ctypes.POINTER(mts_block_t)()
|
|
416
461
|
|
|
@@ -463,11 +508,13 @@ class TensorBlock:
|
|
|
463
508
|
... )
|
|
464
509
|
>>> block.add_gradient("parameter", gradient)
|
|
465
510
|
>>> print(block)
|
|
466
|
-
TensorBlock
|
|
467
|
-
samples
|
|
468
|
-
components
|
|
469
|
-
properties
|
|
470
|
-
|
|
511
|
+
TensorBlock with shape (3, 1, 1)
|
|
512
|
+
samples: [system]
|
|
513
|
+
components: [component]
|
|
514
|
+
properties: [property]
|
|
515
|
+
|
|
516
|
+
gradients:
|
|
517
|
+
parameter => TensorBlock with shape (2, 1, 1)
|
|
471
518
|
"""
|
|
472
519
|
if self._parent is not None:
|
|
473
520
|
raise ValueError(
|
|
@@ -12,41 +12,43 @@ import ctypes
|
|
|
12
12
|
import platform
|
|
13
13
|
from ctypes import CFUNCTYPE, POINTER
|
|
14
14
|
|
|
15
|
+
from ctypes_dlpack import DLDataType, DLDevice, DLManagedTensorVersioned, DLPackVersion
|
|
15
16
|
|
|
16
|
-
class EnumType(type(ctypes.c_int32)):
|
|
17
|
-
def __new__(metacls, name, bases, dict):
|
|
18
|
-
if "_members_" not in dict:
|
|
19
|
-
_members_ = {}
|
|
20
|
-
for key, value in dict.items():
|
|
21
|
-
if not key.startswith("_"):
|
|
22
|
-
_members_[key] = value
|
|
23
17
|
|
|
24
|
-
|
|
18
|
+
class _EnumType(type(ctypes.c_int32)):
|
|
19
|
+
def __new__(metacls, name, bases, namespace):
|
|
20
|
+
if "_members_" not in namespace:
|
|
21
|
+
members = {}
|
|
22
|
+
for key, value in namespace.items():
|
|
23
|
+
if not key.startswith("_"):
|
|
24
|
+
members[key] = value
|
|
25
|
+
namespace["_members_"] = members
|
|
25
26
|
else:
|
|
26
|
-
|
|
27
|
+
members = namespace["_members_"]
|
|
27
28
|
|
|
28
|
-
|
|
29
|
-
|
|
30
|
-
for key, value in cls._members_.items():
|
|
31
|
-
globals()[key] = value
|
|
32
|
-
return cls
|
|
29
|
+
namespace["_reverse_map_"] = {v: k for k, v in members.items()}
|
|
30
|
+
return type(ctypes.c_int32).__new__(metacls, name, bases, namespace)
|
|
33
31
|
|
|
34
32
|
def __repr__(self):
|
|
35
|
-
return "<Enum
|
|
33
|
+
return f"<Enum {self.__name__}>"
|
|
36
34
|
|
|
37
35
|
|
|
38
|
-
class
|
|
36
|
+
class _Enum(ctypes.c_int32, metaclass=_EnumType):
|
|
39
37
|
_members_ = {}
|
|
40
38
|
|
|
41
39
|
def __repr__(self):
|
|
42
|
-
|
|
43
|
-
return f"{self.__class__.__name__}.{
|
|
40
|
+
value_name = self._reverse_map_.get(self.value, str(self.value))
|
|
41
|
+
return f"{self.__class__.__name__}.{value_name}"
|
|
44
42
|
|
|
45
43
|
def __eq__(self, other):
|
|
46
44
|
if isinstance(other, int):
|
|
47
45
|
return self.value == other
|
|
46
|
+
if type(self) is type(other):
|
|
47
|
+
return self.value == other.value
|
|
48
|
+
return NotImplemented
|
|
48
49
|
|
|
49
|
-
|
|
50
|
+
def __hash__(self):
|
|
51
|
+
return hash(self.value)
|
|
50
52
|
|
|
51
53
|
|
|
52
54
|
arch = platform.architecture()[0]
|
|
@@ -57,47 +59,7 @@ elif arch == "64bit":
|
|
|
57
59
|
|
|
58
60
|
|
|
59
61
|
|
|
60
|
-
class
|
|
61
|
-
kDLCPU = 1
|
|
62
|
-
kDLCUDA = 2
|
|
63
|
-
kDLCUDAHost = 3
|
|
64
|
-
kDLOpenCL = 4
|
|
65
|
-
kDLVulkan = 7
|
|
66
|
-
kDLMetal = 8
|
|
67
|
-
kDLVPI = 9
|
|
68
|
-
kDLROCM = 10
|
|
69
|
-
kDLROCMHost = 11
|
|
70
|
-
kDLExtDev = 12
|
|
71
|
-
kDLCUDAManaged = 13
|
|
72
|
-
kDLOneAPI = 14
|
|
73
|
-
kDLWebGPU = 15
|
|
74
|
-
kDLHexagon = 16
|
|
75
|
-
kDLMAIA = 17
|
|
76
|
-
kDLTrn = 18
|
|
77
|
-
|
|
78
|
-
|
|
79
|
-
class DLDataTypeCode(Enum):
|
|
80
|
-
kDLInt = 0
|
|
81
|
-
kDLUInt = 1
|
|
82
|
-
kDLFloat = 2
|
|
83
|
-
kDLOpaqueHandle = 3
|
|
84
|
-
kDLBfloat = 4
|
|
85
|
-
kDLComplex = 5
|
|
86
|
-
kDLBool = 6
|
|
87
|
-
kDLFloat8_e3m4 = 7
|
|
88
|
-
kDLFloat8_e4m3 = 8
|
|
89
|
-
kDLFloat8_e4m3b11fnuz = 9
|
|
90
|
-
kDLFloat8_e4m3fn = 10
|
|
91
|
-
kDLFloat8_e4m3fnuz = 11
|
|
92
|
-
kDLFloat8_e5m2 = 12
|
|
93
|
-
kDLFloat8_e5m2fnuz = 13
|
|
94
|
-
kDLFloat8_e8m0fnu = 14
|
|
95
|
-
kDLFloat6_e2m3fn = 15
|
|
96
|
-
kDLFloat6_e3m2fn = 16
|
|
97
|
-
kDLFloat4_e2m1fn = 17
|
|
98
|
-
|
|
99
|
-
|
|
100
|
-
class mts_status_t(Enum):
|
|
62
|
+
class mts_status_t(_Enum):
|
|
101
63
|
MTS_SUCCESS = 0
|
|
102
64
|
MTS_INVALID_PARAMETER_ERROR = 1
|
|
103
65
|
MTS_IO_ERROR = 2
|
|
@@ -107,30 +69,6 @@ class mts_status_t(Enum):
|
|
|
107
69
|
MTS_INTERNAL_ERROR = 255
|
|
108
70
|
|
|
109
71
|
|
|
110
|
-
class DLPackVersion(ctypes.Structure):
|
|
111
|
-
pass
|
|
112
|
-
|
|
113
|
-
|
|
114
|
-
class DLDevice(ctypes.Structure):
|
|
115
|
-
pass
|
|
116
|
-
|
|
117
|
-
|
|
118
|
-
class DLDataType(ctypes.Structure):
|
|
119
|
-
pass
|
|
120
|
-
|
|
121
|
-
|
|
122
|
-
class DLTensor(ctypes.Structure):
|
|
123
|
-
pass
|
|
124
|
-
|
|
125
|
-
|
|
126
|
-
class DLManagedTensor(ctypes.Structure):
|
|
127
|
-
pass
|
|
128
|
-
|
|
129
|
-
|
|
130
|
-
class DLManagedTensorVersioned(ctypes.Structure):
|
|
131
|
-
pass
|
|
132
|
-
|
|
133
|
-
|
|
134
72
|
class mts_block_t(ctypes.Structure):
|
|
135
73
|
pass
|
|
136
74
|
|
|
@@ -151,56 +89,11 @@ class mts_array_t(ctypes.Structure):
|
|
|
151
89
|
pass
|
|
152
90
|
|
|
153
91
|
|
|
154
|
-
DLPackManagedTensorAllocator = CFUNCTYPE(ctypes.c_int, POINTER(DLTensor), POINTER(POINTER(DLManagedTensorVersioned)), ctypes.c_void_p, CFUNCTYPE(None, ctypes.c_void_p, ctypes.c_char_p, ctypes.c_char_p))
|
|
155
|
-
DLPackManagedTensorFromPyObjectNoSync = CFUNCTYPE(ctypes.c_int, ctypes.c_void_p, POINTER(POINTER(DLManagedTensorVersioned)))
|
|
156
|
-
DLPackDLTensorFromPyObjectNoSync = CFUNCTYPE(ctypes.c_int, ctypes.c_void_p, POINTER(DLTensor))
|
|
157
|
-
DLPackCurrentWorkStream = CFUNCTYPE(ctypes.c_int, DLDeviceType, ctypes.c_int32, POINTER(POINTER(None)))
|
|
158
|
-
DLPackManagedTensorToPyObjectNoSync = CFUNCTYPE(ctypes.c_int, POINTER(DLManagedTensorVersioned), POINTER(POINTER(None)))
|
|
159
92
|
mts_data_origin_t = ctypes.c_uint64
|
|
160
93
|
mts_realloc_buffer_t = CFUNCTYPE(ctypes.c_char_p, ctypes.c_void_p, ctypes.c_char_p, c_uintptr_t)
|
|
161
94
|
mts_create_array_callback_t = CFUNCTYPE(mts_status_t, POINTER(c_uintptr_t), c_uintptr_t, DLDataType, POINTER(mts_array_t))
|
|
162
95
|
|
|
163
96
|
|
|
164
|
-
DLPackVersion._fields_ = [
|
|
165
|
-
("major", ctypes.c_uint32),
|
|
166
|
-
("minor", ctypes.c_uint32),
|
|
167
|
-
]
|
|
168
|
-
|
|
169
|
-
DLDevice._fields_ = [
|
|
170
|
-
("device_type", DLDeviceType),
|
|
171
|
-
("device_id", ctypes.c_int32),
|
|
172
|
-
]
|
|
173
|
-
|
|
174
|
-
DLDataType._fields_ = [
|
|
175
|
-
("code", ctypes.c_uint8),
|
|
176
|
-
("bits", ctypes.c_uint8),
|
|
177
|
-
("lanes", ctypes.c_uint16),
|
|
178
|
-
]
|
|
179
|
-
|
|
180
|
-
DLTensor._fields_ = [
|
|
181
|
-
("data", ctypes.c_void_p),
|
|
182
|
-
("device", DLDevice),
|
|
183
|
-
("ndim", ctypes.c_int32),
|
|
184
|
-
("dtype", DLDataType),
|
|
185
|
-
("shape", POINTER(ctypes.c_int64)),
|
|
186
|
-
("strides", POINTER(ctypes.c_int64)),
|
|
187
|
-
("byte_offset", ctypes.c_uint64),
|
|
188
|
-
]
|
|
189
|
-
|
|
190
|
-
DLManagedTensor._fields_ = [
|
|
191
|
-
("dl_tensor", DLTensor),
|
|
192
|
-
("manager_ctx", ctypes.c_void_p),
|
|
193
|
-
("deleter", CFUNCTYPE(None, POINTER(DLManagedTensor))),
|
|
194
|
-
]
|
|
195
|
-
|
|
196
|
-
DLManagedTensorVersioned._fields_ = [
|
|
197
|
-
("version", DLPackVersion),
|
|
198
|
-
("manager_ctx", ctypes.c_void_p),
|
|
199
|
-
("deleter", CFUNCTYPE(None, POINTER(DLManagedTensorVersioned))),
|
|
200
|
-
("flags", ctypes.c_uint64),
|
|
201
|
-
("dl_tensor", DLTensor),
|
|
202
|
-
]
|
|
203
|
-
|
|
204
97
|
mts_data_movement_t._fields_ = [
|
|
205
98
|
("sample_in", c_uintptr_t),
|
|
206
99
|
("sample_out", c_uintptr_t),
|
|
@@ -2,29 +2,23 @@ import ctypes
|
|
|
2
2
|
from typing import NewType, Union
|
|
3
3
|
|
|
4
4
|
import numpy as np
|
|
5
|
-
|
|
6
|
-
from .._c_api import (
|
|
5
|
+
from ctypes_dlpack import (
|
|
7
6
|
DLDataType,
|
|
8
7
|
DLDevice,
|
|
9
8
|
DLDeviceType,
|
|
10
9
|
DLManagedTensorVersioned,
|
|
10
|
+
DLPackArray,
|
|
11
11
|
DLPackVersion,
|
|
12
|
+
array_as_dlpack,
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
from .._c_api import (
|
|
12
16
|
c_uintptr_t,
|
|
13
17
|
mts_array_t,
|
|
14
18
|
mts_data_movement_t,
|
|
15
19
|
mts_data_origin_t,
|
|
16
20
|
)
|
|
17
21
|
from .._status import catch_exceptions, check_status
|
|
18
|
-
from ._dlpack import (
|
|
19
|
-
DLPACK_NAME,
|
|
20
|
-
DLPACK_VERSIONED_NAME,
|
|
21
|
-
PYTHON_API,
|
|
22
|
-
USED_DLPACK_NAME,
|
|
23
|
-
USED_DLPACK_VERSIONED_NAME,
|
|
24
|
-
DLManagedTensor,
|
|
25
|
-
DLPackArray,
|
|
26
|
-
wrap_unversioned_as_versioned,
|
|
27
|
-
)
|
|
28
22
|
|
|
29
23
|
|
|
30
24
|
try:
|
|
@@ -37,9 +31,9 @@ except ImportError:
|
|
|
37
31
|
if HAS_TORCH:
|
|
38
32
|
# This NewType is only used for typechecking and documentation purposes. If you are
|
|
39
33
|
# trying to add support for new array types, see `data.array.ArrayWrapper` instead.
|
|
40
|
-
Array = NewType("Array", Union[np.ndarray, torch.Tensor])
|
|
34
|
+
Array = NewType("Array", Union[np.ndarray, torch.Tensor, mts_array_t])
|
|
41
35
|
else:
|
|
42
|
-
Array = NewType("Array", np.ndarray)
|
|
36
|
+
Array = NewType("Array", Union[np.ndarray, mts_array_t])
|
|
43
37
|
|
|
44
38
|
Array.__doc__ = """
|
|
45
39
|
An ``Array`` contains the actual data stored in a :py:class:`metatensor.TensorBlock`.
|
|
@@ -500,107 +494,21 @@ def _mts_array_move_data(
|
|
|
500
494
|
@catch_exceptions
|
|
501
495
|
def _mts_array_as_dlpack(this, dl_managed_tensor_ptr_ptr, device, stream, max_version):
|
|
502
496
|
"""
|
|
503
|
-
Implementation of
|
|
497
|
+
Implementation of ``mts_array_t.as_dlpack``.
|
|
504
498
|
|
|
505
|
-
This function calls the array's __dlpack__ method, gets the PyCapsule,
|
|
506
|
-
|
|
507
|
-
ownership to the C-API caller.
|
|
508
|
-
|
|
509
|
-
When the PyCapsule contains a `dltensor_versioned`, we just need to rename
|
|
510
|
-
the capsule to `used_dltensor_versioned` after giving the data to the C API.
|
|
511
|
-
|
|
512
|
-
When the PyCapsule contains a deprecated `dltensor`, we re-wrap the data
|
|
513
|
-
inside `DLManagedTensorVersioned`.
|
|
499
|
+
This function calls the array's __dlpack__ method, gets the PyCapsule, extracts the
|
|
500
|
+
raw pointer, and transfers ownership to the C-API caller.
|
|
514
501
|
"""
|
|
515
502
|
wrapper = _KNOWN_ARRAY_WRAPPERS[this]
|
|
516
|
-
array = wrapper.array
|
|
517
|
-
|
|
518
|
-
dl_device = (device.device_type.value, device.device_id)
|
|
519
|
-
stream = stream.contents if stream else None
|
|
520
|
-
max_version = (max_version.major, max_version.minor)
|
|
521
503
|
|
|
522
|
-
|
|
523
|
-
|
|
524
|
-
|
|
525
|
-
# Try requesting versioned DLPack
|
|
526
|
-
capsule = array.__dlpack__(
|
|
527
|
-
stream=stream, max_version=max_version, dl_device=dl_device
|
|
528
|
-
)
|
|
529
|
-
except Exception as _:
|
|
530
|
-
# Fallback to legacy signatures. Each fallback drops one parameter.
|
|
531
|
-
try:
|
|
532
|
-
capsule = array.__dlpack__(stream=stream, dl_device=dl_device)
|
|
533
|
-
except Exception:
|
|
534
|
-
try:
|
|
535
|
-
capsule = array.__dlpack__(dl_device=dl_device)
|
|
536
|
-
except Exception:
|
|
537
|
-
capsule = array.__dlpack__()
|
|
538
|
-
|
|
539
|
-
capsule_name = PYTHON_API.PyCapsule_GetName(capsule)
|
|
540
|
-
|
|
541
|
-
# Versioned Capsule
|
|
542
|
-
if capsule_name == DLPACK_VERSIONED_NAME:
|
|
543
|
-
pointer = PYTHON_API.PyCapsule_GetPointer(capsule, DLPACK_VERSIONED_NAME)
|
|
544
|
-
if not pointer:
|
|
545
|
-
raise
|
|
546
|
-
|
|
547
|
-
versioned_ptr = ctypes.cast(pointer, ctypes.POINTER(DLManagedTensorVersioned))
|
|
548
|
-
actual_device = versioned_ptr.contents.dl_tensor.device
|
|
549
|
-
if (
|
|
550
|
-
actual_device.device_type != device.device_type
|
|
551
|
-
or actual_device.device_id != device.device_id
|
|
552
|
-
):
|
|
553
|
-
raise ValueError(
|
|
554
|
-
f"DLPack device mismatch: expected type={device.device_type} "
|
|
555
|
-
f"id={device.device_id}, got type={actual_device.device_type} "
|
|
556
|
-
f"id={actual_device.device_id}"
|
|
557
|
-
)
|
|
558
|
-
|
|
559
|
-
status = PYTHON_API.PyCapsule_SetName(capsule, USED_DLPACK_VERSIONED_NAME)
|
|
560
|
-
if status != 0:
|
|
561
|
-
raise
|
|
562
|
-
|
|
563
|
-
dl_managed_tensor_ptr_ptr[0] = versioned_ptr
|
|
564
|
-
return
|
|
565
|
-
|
|
566
|
-
# Legacy Capsule
|
|
567
|
-
if capsule_name == DLPACK_NAME:
|
|
568
|
-
pointer = PYTHON_API.PyCapsule_GetPointer(capsule, DLPACK_NAME)
|
|
569
|
-
if not pointer:
|
|
570
|
-
raise
|
|
571
|
-
|
|
572
|
-
unversioned_ptr = ctypes.cast(pointer, ctypes.POINTER(DLManagedTensor))
|
|
573
|
-
|
|
574
|
-
actual_device = unversioned_ptr.contents.dl_tensor.device
|
|
575
|
-
if (
|
|
576
|
-
actual_device.device_type != device.device_type
|
|
577
|
-
or actual_device.device_id != device.device_id
|
|
578
|
-
):
|
|
579
|
-
raise ValueError(
|
|
580
|
-
f"DLPack device mismatch: expected type {device.device_type}"
|
|
581
|
-
f" id {device.device_id}, "
|
|
582
|
-
f"got type {actual_device.device_type}"
|
|
583
|
-
f" id {actual_device.device_id}"
|
|
584
|
-
)
|
|
585
|
-
|
|
586
|
-
# We must do this to ensure that the deletion of the PyCapsule does not run the
|
|
587
|
-
# deleter and cause double-free issues. Instead we call legacy deleter
|
|
588
|
-
# explicitly as part of the versioned tensor's deleter.
|
|
589
|
-
status = PYTHON_API.PyCapsule_SetName(capsule, USED_DLPACK_NAME)
|
|
590
|
-
if status != 0:
|
|
591
|
-
raise
|
|
592
|
-
|
|
593
|
-
# Wrap it
|
|
594
|
-
versioned_ptr = wrap_unversioned_as_versioned(unversioned_ptr)
|
|
595
|
-
dl_managed_tensor_ptr_ptr[0] = versioned_ptr
|
|
596
|
-
return
|
|
597
|
-
|
|
598
|
-
raise ValueError(
|
|
599
|
-
"Unexpected DLPack capsule name:"
|
|
600
|
-
f" '{capsule_name.decode() if capsule_name else 'NULL'}'. "
|
|
601
|
-
f"Expected '{DLPACK_VERSIONED_NAME.decode()}' or '{DLPACK_NAME.decode()}'"
|
|
504
|
+
stream = None if not stream else stream.contents
|
|
505
|
+
versioned_ptr = array_as_dlpack(
|
|
506
|
+
wrapper.array, dl_device=device, stream=stream, max_version=max_version
|
|
602
507
|
)
|
|
603
508
|
|
|
509
|
+
dl_managed_tensor_ptr_ptr[0] = versioned_ptr
|
|
510
|
+
return
|
|
511
|
+
|
|
604
512
|
|
|
605
513
|
@catch_exceptions
|
|
606
514
|
def _mts_array_from_dlpack(this, dl_tensor, new_array):
|
|
@@ -2,12 +2,15 @@ import ctypes
|
|
|
2
2
|
from typing import Any
|
|
3
3
|
|
|
4
4
|
import numpy as np
|
|
5
|
-
|
|
6
|
-
from .._c_api import (
|
|
5
|
+
from ctypes_dlpack import (
|
|
7
6
|
DLDevice,
|
|
8
7
|
DLDeviceType,
|
|
9
8
|
DLManagedTensorVersioned,
|
|
9
|
+
DLPackArray,
|
|
10
10
|
DLPackVersion,
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
from .._c_api import (
|
|
11
14
|
c_uintptr_t,
|
|
12
15
|
mts_array_t,
|
|
13
16
|
mts_data_origin_t,
|
|
@@ -20,7 +23,6 @@ from ._array import (
|
|
|
20
23
|
_origin_pytorch,
|
|
21
24
|
_register_origin,
|
|
22
25
|
)
|
|
23
|
-
from ._dlpack import DLPackArray, wrap_versioned_as_unversioned
|
|
24
26
|
|
|
25
27
|
|
|
26
28
|
_ADDITIONAL_ORIGINS = {}
|
|
@@ -235,13 +237,7 @@ class ExternalCudaArray:
|
|
|
235
237
|
pass
|
|
236
238
|
|
|
237
239
|
if tensor is None:
|
|
238
|
-
|
|
239
|
-
tensor = torch.from_dlpack(dlpack_array)
|
|
240
|
-
except RuntimeError:
|
|
241
|
-
# Older PyTorch (< 2.4) doesn't understand versioned
|
|
242
|
-
# DLPack capsules. Fall back to an unversioned capsule.
|
|
243
|
-
unversioned = wrap_versioned_as_unversioned(dl_managed_ptr)
|
|
244
|
-
tensor = torch.from_dlpack(DLPackArray(unversioned))
|
|
240
|
+
tensor = torch.from_dlpack(dlpack_array)
|
|
245
241
|
|
|
246
242
|
# keep a reference to the parent object to prevent it from being
|
|
247
243
|
# garbage-collected while the tensor is alive
|