torchsparsegradutils 0.1.0__tar.gz → 0.1.2__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 (54) hide show
  1. {torchsparsegradutils-0.1.0/torchsparsegradutils.egg-info → torchsparsegradutils-0.1.2}/PKG-INFO +2 -2
  2. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/README.md +1 -1
  3. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/setup.py +4 -4
  4. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/distributions/sparse_multivariate_normal.py +1 -0
  5. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/encoders/pairwise_voxel_encoder.py +29 -5
  6. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_encoders.py +22 -3
  7. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2/torchsparsegradutils.egg-info}/PKG-INFO +2 -2
  8. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/LICENSE +0 -0
  9. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/MANIFEST.in +0 -0
  10. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/pyproject.toml +0 -0
  11. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/setup.cfg +0 -0
  12. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/__init__.py +0 -0
  13. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/cupy/__init__.py +0 -0
  14. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/cupy/cupy_bindings.py +0 -0
  15. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/cupy/cupy_sparse_solve.py +0 -0
  16. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/distributions/__init__.py +0 -0
  17. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/distributions/constraints.py +0 -0
  18. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/encoders/__init__.py +0 -0
  19. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/jax/__init__.py +0 -0
  20. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/jax/jax_bindings.py +0 -0
  21. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/jax/jax_sparse_solve.py +0 -0
  22. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/sparse_lstsq.py +0 -0
  23. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/sparse_matmul.py +0 -0
  24. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/sparse_solve.py +0 -0
  25. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/__init__.py +0 -0
  26. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_bicgstab.py +0 -0
  27. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_cupy_bindings.py +0 -0
  28. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_cupy_sparse_solve.py +0 -0
  29. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_distributions.py +0 -0
  30. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_jax_bindings.py +0 -0
  31. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_jax_sparse_solve.py +0 -0
  32. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_linear_cg.py +0 -0
  33. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_lsmr.py +0 -0
  34. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_minres.py +0 -0
  35. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_params/czyx_shifts.yaml +0 -0
  36. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_params/pairwise_coo_indices.yaml +0 -0
  37. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_params/xyz_coords.yaml +0 -0
  38. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_random.py +0 -0
  39. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_sparse_lstsq.py +0 -0
  40. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_sparse_matmul.py +0 -0
  41. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_sparse_solve.py +0 -0
  42. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/tests/test_utils.py +0 -0
  43. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/utils/__init__.py +0 -0
  44. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/utils/bicgstab.py +0 -0
  45. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/utils/linear_cg.py +0 -0
  46. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/utils/lsmr.py +0 -0
  47. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/utils/minres.py +0 -0
  48. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/utils/random_sparse.py +0 -0
  49. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils/utils/utils.py +0 -0
  50. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils.egg-info/SOURCES.txt +0 -0
  51. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils.egg-info/dependency_links.txt +0 -0
  52. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils.egg-info/not-zip-safe +0 -0
  53. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils.egg-info/requires.txt +0 -0
  54. {torchsparsegradutils-0.1.0 → torchsparsegradutils-0.1.2}/torchsparsegradutils.egg-info/top_level.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: torchsparsegradutils
3
- Version: 0.1.0
3
+ Version: 0.1.2
4
4
  Summary: A collection of utility functions to work with PyTorch sparse tensors
5
5
  Home-page: https://github.com/cai4cai/torchsparsegradutils
6
6
  Author: CAI4CAI research group
@@ -43,7 +43,7 @@ Things that are missing may be listed as [issues](https://github.com/cai4cai/tor
43
43
  ## Installation
44
44
  The provided package can be installed using:
45
45
 
46
- `pip install torchsparsegradutils` (TODO)
46
+ `pip install torchsparsegradutils`
47
47
 
48
48
  or
49
49
 
@@ -24,7 +24,7 @@ Things that are missing may be listed as [issues](https://github.com/cai4cai/tor
24
24
  ## Installation
25
25
  The provided package can be installed using:
26
26
 
27
- `pip install torchsparsegradutils` (TODO)
27
+ `pip install torchsparsegradutils`
28
28
 
29
29
  or
30
30
 
@@ -8,7 +8,7 @@ def readme():
8
8
 
9
9
  setuptools.setup(
10
10
  name="torchsparsegradutils",
11
- version="0.1.0",
11
+ version="0.1.2",
12
12
  description="A collection of utility functions to work with PyTorch sparse tensors",
13
13
  long_description=readme(),
14
14
  long_description_content_type="text/markdown",
@@ -19,7 +19,7 @@ setuptools.setup(
19
19
  "Programming Language :: Python :: 3.9",
20
20
  "Programming Language :: Python :: 3.10",
21
21
  ],
22
- python_requires='>=3.8, <3.11',
22
+ python_requires=">=3.8, <3.11",
23
23
  keywords="sparse torch utility",
24
24
  url="https://github.com/cai4cai/torchsparsegradutils",
25
25
  author="CAI4CAI research group",
@@ -29,8 +29,8 @@ setuptools.setup(
29
29
  install_requires=[
30
30
  "torch>=1.13",
31
31
  ],
32
- setup_requires=['pytest-runner'],
33
- tests_require=['pytest'],
32
+ setup_requires=["pytest-runner"],
33
+ tests_require=["pytest"],
34
34
  test_suite="tests",
35
35
  extras_require={
36
36
  "extras": ["jax", "cupy"],
@@ -76,6 +76,7 @@ class SparseMultivariateNormal(Distribution):
76
76
  in either torch.sparse_coo or torch.sparse_csr layout
77
77
  """
78
78
 
79
+ arg_constraints = {}
79
80
  # TODO: add in constraints
80
81
  # arg_constraints = {'loc': constraints.real_vector,
81
82
  # 'diag': constraints.independent(constraints.positive, 1),
@@ -257,7 +257,7 @@ def calc_pariwise_coo_indices(
257
257
  return indices
258
258
 
259
259
 
260
- class PairwiseVoxelEncoder:
260
+ class PairwiseVoxelEncoder(torch.nn.Module):
261
261
  """
262
262
  A class for encoding pairwise spatial local neighbourhoods and channel based voxel relations
263
263
  into a sparse tensor of either COO or CSR format.
@@ -270,6 +270,8 @@ class PairwiseVoxelEncoder:
270
270
  Additionally, diagonal entries can be included with the diag flag.
271
271
  The output matrix can be restricted to upper or lower triangular with the upper flag, if
272
272
  symmetric relationships are assumed, such as distance or correlation.
273
+ The indices are stored on the device specified by the device argument,
274
+ these indices can be sent to another device using the to() method.
273
275
 
274
276
  The sparse tensor is returned in the `__call__` method, which takes a tensor of values
275
277
  with shape [(B), N, C, H, D, W] and returns a sparse tensor of shape [(B), S, S]
@@ -284,6 +286,7 @@ class PairwiseVoxelEncoder:
284
286
  Some of the values will be trimmed from the 3D volume edges, as values are not allowed
285
287
  to wrap around the edges of the volume of spatial volume.
286
288
 
289
+ The sparse tensor is returned on the same device as the input values to the `__call__` method.
287
290
 
288
291
  Args:
289
292
  radius (float): The maximum distance from the origin within which the spatial
@@ -309,7 +312,8 @@ class PairwiseVoxelEncoder:
309
312
  indices_dtype (torch.dtype, optional): The data type of the output indices.
310
313
  Must be either torch.int32 or torch.int64.
311
314
  Default is torch.int64.
312
- device (torch.device, optional): Device assigned to generate sparse tensor.
315
+ device (torch.device, optional): Device assigned to store sparse tensor indices
316
+ at initialisation.
313
317
  Defaults to torch.device("cpu").
314
318
 
315
319
 
@@ -333,6 +337,8 @@ class PairwiseVoxelEncoder:
333
337
  indices_dtype: torch.dtype = torch.int64,
334
338
  device: torch.device = torch.device("cpu"),
335
339
  ):
340
+ super().__init__()
341
+
336
342
  if not ((len(volume_shape) == 4) and all(isinstance(dim, int) and dim > 0 for dim in volume_shape)):
337
343
  raise ValueError("`volume_shape` must be a 4D tuple of positive integers, representing [C, H, D, W]")
338
344
 
@@ -348,7 +354,6 @@ class PairwiseVoxelEncoder:
348
354
  self.channel_voxel_relation = channel_voxel_relation
349
355
  self.layout = layout
350
356
  self.indices_dtype = indices_dtype
351
- self.device = device
352
357
 
353
358
  self.volume_numel = reduce(mul, volume_shape)
354
359
 
@@ -369,6 +374,23 @@ class PairwiseVoxelEncoder:
369
374
  else:
370
375
  raise ValueError("layout must be either torch.sparse_coo or torch.sparse_csr")
371
376
 
377
+ def _apply(self, fn):
378
+ # Applying the function to the desired attributes
379
+ # This has been implemented to allow using the .to() method
380
+ for attr in ["indices", "csr_permutation", "crow_indices", "col_indices"]:
381
+ tensor = getattr(self, attr, None)
382
+ if tensor is not None:
383
+ setattr(self, attr, fn(tensor))
384
+
385
+ return self
386
+
387
+ @property
388
+ def device(self):
389
+ if self.layout == torch.sparse_coo:
390
+ return self.indices.device
391
+ elif self.layout == torch.sparse_csr:
392
+ return self.crow_indices.device
393
+
372
394
  def _calc_values(self, values: torch.Tensor) -> torch.Tensor:
373
395
  """
374
396
  Calculate the values for the sparse tensor based on the input values and offsets.
@@ -398,6 +420,8 @@ class PairwiseVoxelEncoder:
398
420
  boundaries of the input volume are trimmed to ensure that the output tensor doesn't
399
421
  contain any relationships that wrap around the spatial volume described by this encoder.
400
422
 
423
+ The output sparse tensor will be returned on the same device as the input values tensor.
424
+
401
425
  Args:
402
426
  values (torch.Tensor): Input tensor of values with shape [(B), N, C, H, D, W]
403
427
  where B is an optional batch dimension, N is the number of offsets
@@ -407,7 +431,7 @@ class PairwiseVoxelEncoder:
407
431
 
408
432
  Returns:
409
433
  torch.Tensor: Output tensor in either COO or CSR format, with shape [(B), S, S]
410
- where is C*H*D*W.
434
+ where S is C*H*D*W.
411
435
 
412
436
  Raises:
413
437
  ValueError: If the shape of 'values' is not a 5D or 6D tensor.
@@ -452,7 +476,7 @@ class PairwiseVoxelEncoder:
452
476
  else:
453
477
  sparse_dim_indices = self.indices.repeat(1, batch_size)
454
478
  batch_dim_indices = (
455
- torch.arange(batch_size, dtype=self.indices.dtype, device=self.device)
479
+ torch.arange(batch_size, dtype=self.indices.dtype, device=self.indices.device)
456
480
  .repeat_interleave(self.indices.shape[-1])
457
481
  .unsqueeze(0)
458
482
  )
@@ -157,7 +157,7 @@ def test_trim(tensor_nd, offsets, expected_output_slice):
157
157
  current_dir = Path(os.path.abspath(os.path.dirname(__file__)))
158
158
 
159
159
  # Construct the path to the yaml file:
160
- yaml_file = current_dir / 'test_params' / 'xyz_coords.yaml'
160
+ yaml_file = current_dir / "test_params" / "xyz_coords.yaml"
161
161
 
162
162
  with open(yaml_file) as f: # load test cases from file
163
163
  coord_test_cases = yaml.safe_load(f)
@@ -179,7 +179,7 @@ def test_gen_coords(radius, expected_coords):
179
179
 
180
180
  # Test neighbourgood offset generation:
181
181
 
182
- yaml_file = current_dir / 'test_params' / 'czyx_shifts.yaml'
182
+ yaml_file = current_dir / "test_params" / "czyx_shifts.yaml"
183
183
 
184
184
  with open(yaml_file) as f: # load test cases from file
185
185
  shift_test_cases = yaml.safe_load(f)
@@ -261,7 +261,7 @@ def test_pairwise_coo_indices_unique(radius, volume_shape, diag, upper, channel_
261
261
 
262
262
  # Test the indices generated are as expected for a simple (3, 2, 2, 2) volume:
263
263
 
264
- yaml_file = current_dir / 'test_params' / 'pairwise_coo_indices.yaml'
264
+ yaml_file = current_dir / "test_params" / "pairwise_coo_indices.yaml"
265
265
 
266
266
  with open(yaml_file) as f: # load test cases from file to avoid a massive mess
267
267
  data = yaml.safe_load(f)
@@ -406,6 +406,25 @@ def test_PVE_dtype_device(batch_size, layout, indices_dtype, values_dtype, devic
406
406
  assert sparse_matrix.col_indices().device.type == device.type
407
407
 
408
408
 
409
+ def test_PVE_to_device(layout, device):
410
+ encoder = PairwiseVoxelEncoder(
411
+ radius=1.0,
412
+ volume_shape=(5, 5, 5, 5),
413
+ layout=layout,
414
+ indices_dtype=torch.int64,
415
+ device=torch.device("cpu"),
416
+ )
417
+ encoder.to(device)
418
+
419
+ if layout == torch.sparse_coo:
420
+ assert encoder.indices.device.type == device.type
421
+
422
+ elif layout == torch.sparse_csr:
423
+ assert encoder.crow_indices.device.type == device.type
424
+ assert encoder.col_indices.device.type == device.type
425
+ assert encoder.csr_permutation.device.type == device.type
426
+
427
+
409
428
  # Test based on expected indices in pairwise_coo_indices.yaml
410
429
  @pytest.mark.parametrize(
411
430
  "radius, volume_shape, diag, upper, channel_relation, expected_indices",
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: torchsparsegradutils
3
- Version: 0.1.0
3
+ Version: 0.1.2
4
4
  Summary: A collection of utility functions to work with PyTorch sparse tensors
5
5
  Home-page: https://github.com/cai4cai/torchsparsegradutils
6
6
  Author: CAI4CAI research group
@@ -43,7 +43,7 @@ Things that are missing may be listed as [issues](https://github.com/cai4cai/tor
43
43
  ## Installation
44
44
  The provided package can be installed using:
45
45
 
46
- `pip install torchsparsegradutils` (TODO)
46
+ `pip install torchsparsegradutils`
47
47
 
48
48
  or
49
49