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.
Files changed (42) hide show
  1. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/CMakeLists.txt +1 -1
  2. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/PKG-INFO +2 -1
  3. metatensor_core-0.2.4/git_version_info +2 -0
  4. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_block.py +90 -43
  5. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_c_api.py +22 -129
  6. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_data/_array.py +17 -109
  7. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_data/_extract.py +6 -10
  8. metatensor_core-0.2.4/metatensor/_html.py +425 -0
  9. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_labels.py +126 -84
  10. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_status.py +1 -1
  11. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_tensor.py +170 -10
  12. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/io/_block.py +1 -2
  13. metatensor_core-0.2.4/metatensor-core-cxx-0.2.4.tar.gz +0 -0
  14. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor_core.egg-info/PKG-INFO +2 -1
  15. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor_core.egg-info/SOURCES.txt +2 -2
  16. metatensor_core-0.2.4/metatensor_core.egg-info/requires.txt +2 -0
  17. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/pyproject.toml +1 -1
  18. metatensor_core-0.2.2/git_version_info +0 -2
  19. metatensor_core-0.2.2/metatensor/_data/_dlpack.py +0 -293
  20. metatensor_core-0.2.2/metatensor-core-cxx-0.2.2.tar.gz +0 -0
  21. metatensor_core-0.2.2/metatensor_core.egg-info/requires.txt +0 -1
  22. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/AUTHORS +0 -0
  23. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/LICENSE +0 -0
  24. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/MANIFEST.in +0 -0
  25. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/README.rst +0 -0
  26. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/__init__.py +0 -0
  27. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_c_lib.py +0 -0
  28. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_data/__init__.py +0 -0
  29. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/_version.py +0 -0
  30. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/io/__init__.py +0 -0
  31. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/io/_labels.py +0 -0
  32. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/io/_tensor.py +0 -0
  33. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/io/_utils.py +0 -0
  34. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/learn.py +0 -0
  35. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/operations.py +0 -0
  36. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/torch.py +0 -0
  37. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor/utils.py +0 -0
  38. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor_core.egg-info/dependency_links.txt +0 -0
  39. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor_core.egg-info/not-zip-safe +0 -0
  40. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/metatensor_core.egg-info/top_level.txt +0 -0
  41. {metatensor_core-0.2.2 → metatensor_core-0.2.4}/setup.cfg +0 -0
  42. {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.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.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
 
@@ -0,0 +1,2 @@
1
+ 0
2
+ git.417f531
@@ -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 (2): ['samples']
53
- components (): []
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 = f"Gradient TensorBlock ('{'/'.join(self._gradient_parameters)}')\n"
259
- else:
260
- s = "TensorBlock\n"
261
- s += f" samples ({len(self.samples)}): {str(list(self.samples.names))}"
262
- s += "\n"
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 += "None"
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
- Gradient TensorBlock ('positions')
402
- samples (2): ['sample', 'atom']
403
- components (3, 1): ['direction', 'component']
404
- properties (5): ['property']
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
- Gradient TensorBlock ('cell')
410
- samples (2): ['sample']
411
- components (3, 3, 1): ['direction_1', 'direction_2', 'component']
412
- properties (5): ['property']
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 (3): ['system']
468
- components (1): ['component']
469
- properties (1): ['property']
470
- gradients: ['parameter']
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
- dict["_members_"] = _members_
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
- _members_ = dict["_members_"]
27
+ members = namespace["_members_"]
27
28
 
28
- dict["_reverse_map_"] = {v: k for k, v in _members_.items()}
29
- cls = type(ctypes.c_int32).__new__(metacls, name, bases, dict)
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 %s>" % self.__name__
33
+ return f"<Enum {self.__name__}>"
36
34
 
37
35
 
38
- class Enum(ctypes.c_int32, metaclass=EnumType):
36
+ class _Enum(ctypes.c_int32, metaclass=_EnumType):
39
37
  _members_ = {}
40
38
 
41
39
  def __repr__(self):
42
- value_str = self._reverse_map_.get(self.value, str(self.value))
43
- return f"{self.__class__.__name__}.{value_str}"
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
- return type(self) is type(other) and self.value == other.value
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 DLDeviceType(Enum):
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 `mts_array_t.as_dlpack`.
497
+ Implementation of ``mts_array_t.as_dlpack``.
504
498
 
505
- This function calls the array's __dlpack__ method, gets the PyCapsule,
506
- extracts the raw pointer, and transfers
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
- capsule = None
523
-
524
- try:
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
- try:
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