multipers 2.4.1__tar.gz → 2.4.2b1__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 (237) hide show
  1. {multipers-2.4.1 → multipers-2.4.2b1}/PKG-INFO +1 -1
  2. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/_signed_measure_meta.py +2 -2
  3. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/_slicer_meta.py +5 -4
  4. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/array_api/__init__.py +1 -1
  5. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/array_api/numpy.py +34 -0
  6. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/array_api/torch.py +33 -1
  7. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/filtrations/density.py +57 -54
  8. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/function_rips.pyx +0 -11
  9. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/grids.pyx +180 -88
  10. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Slicer.h +0 -1
  11. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/io.pyx +3 -3
  12. multipers-2.4.1/multipers/ml/point_clouds.py → multipers-2.4.2b1/multipers/ml/filtered_complex.py +142 -101
  13. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/ml/mma.py +10 -1
  14. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/ml/signed_measures.py +109 -80
  15. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/ml/sliced_wasserstein.py +157 -29
  16. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/mma_structures.pxd +2 -1
  17. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/mma_structures.pyx +18 -4
  18. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/mma_structures.pyx.tp +9 -2
  19. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multiparameter_module_approximation/approximation.h +197 -157
  20. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multiparameter_module_approximation.pyx +4 -2
  21. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/ops.pyx +14 -9
  22. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/plots.py +125 -55
  23. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/simplex_tree_multi.pyx +16 -16
  24. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/simplex_tree_multi.pyx.tp +2 -2
  25. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/slicer.pxd +22 -22
  26. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/slicer.pxd.tp +3 -3
  27. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/slicer.pyx +218 -146
  28. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/slicer.pyx.tp +6 -4
  29. {multipers-2.4.1 → multipers-2.4.2b1}/multipers.egg-info/PKG-INFO +1 -1
  30. {multipers-2.4.1 → multipers-2.4.2b1}/multipers.egg-info/SOURCES.txt +1 -4
  31. {multipers-2.4.1 → multipers-2.4.2b1}/pyproject.toml +1 -1
  32. {multipers-2.4.1 → multipers-2.4.2b1}/setup.py +1 -1
  33. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_grids.py +75 -4
  34. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_mma.py +0 -2
  35. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_parallel.py +0 -2
  36. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_point_clouds.py +13 -7
  37. multipers-2.4.1/multipers/torch/__init__.py +0 -1
  38. multipers-2.4.1/multipers/torch/diff_grids.py +0 -240
  39. multipers-2.4.1/multipers/torch/rips_density.py +0 -310
  40. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/aida.cpp +0 -0
  41. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/birth_death.cpp +0 -0
  42. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/brute_force_mpm_decomposition.cpp +0 -0
  43. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/generate_decompositions.cpp +0 -0
  44. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/include/aida_interface.hpp +0 -0
  45. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/include/config.hpp +0 -0
  46. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/include/option_parser.hpp +0 -0
  47. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/include/types.hpp +0 -0
  48. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/making_examples.cpp +0 -0
  49. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/minimize_pres.cpp +0 -0
  50. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/mpfree_clone.cpp +0 -0
  51. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/presentation_to_quiver.cpp +0 -0
  52. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/presentation_to_quiver_new.cpp +0 -0
  53. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/resolution.cpp +0 -0
  54. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/aida_decompose.cpp +0 -0
  55. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/aida_decompose.hpp +0 -0
  56. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/aida_functions.cpp +0 -0
  57. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/aida_functions.hpp +0 -0
  58. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/aida_helpers.cpp +0 -0
  59. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/aida_helpers.hpp +0 -0
  60. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/aida_interface.cpp +0 -0
  61. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/block.cpp +0 -0
  62. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/block.hpp +0 -0
  63. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/config.cpp +0 -0
  64. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/src/option_parser.cpp +0 -0
  65. {multipers-2.4.1 → multipers-2.4.2b1}/AIDA/vectorspace_decompositions.cpp +0 -0
  66. {multipers-2.4.1 → multipers-2.4.2b1}/LICENSE +0 -0
  67. {multipers-2.4.1 → multipers-2.4.2b1}/MANIFEST.in +0 -0
  68. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/bitset_algebra.hpp +0 -0
  69. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/column_types.hpp +0 -0
  70. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/dense_matrix.hpp +0 -0
  71. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/draw_hf.hpp +0 -0
  72. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/general.hpp +0 -0
  73. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/graded_linalg.hpp +0 -0
  74. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/graded_matrix.hpp +0 -0
  75. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/grid_scheduler.hpp +0 -0
  76. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/homomorphisms.hpp +0 -0
  77. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/matrix_base.hpp +0 -0
  78. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/modules.hpp +0 -0
  79. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/orders_and_graphs.hpp +0 -0
  80. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/r2graded_matrix.hpp +0 -0
  81. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/r3graded_matrix.hpp +0 -0
  82. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/sparse_matrix.hpp +0 -0
  83. {multipers-2.4.1 → multipers-2.4.2b1}/Persistence-Algebra/include/grlina/to_quiver.hpp +0 -0
  84. {multipers-2.4.1 → multipers-2.4.2b1}/README.md +0 -0
  85. {multipers-2.4.1 → multipers-2.4.2b1}/_tempita_grid_gen.py +0 -0
  86. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/__init__.py +0 -0
  87. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/data/MOL2.py +0 -0
  88. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/data/UCR.py +0 -0
  89. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/data/__init__.py +0 -0
  90. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/data/graphs.py +0 -0
  91. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/data/immuno_regions.py +0 -0
  92. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/data/minimal_presentation_to_st_bf.py +0 -0
  93. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/data/pytorch2simplextree.py +0 -0
  94. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/data/shape3d.py +0 -0
  95. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/data/synthetic.py +0 -0
  96. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/distances.py +0 -0
  97. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/filtration_conversions.pxd +0 -0
  98. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/filtration_conversions.pxd.tp +0 -0
  99. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/filtrations/__init__.py +0 -0
  100. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/filtrations/filtrations.py +0 -0
  101. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/filtrations.pxd +0 -0
  102. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/filtrations.pxd.tp +0 -0
  103. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/Persistence_slices_interface.h +0 -0
  104. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/Simplex_tree_interface.h +0 -0
  105. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/Simplex_tree_multi_interface.h +0 -0
  106. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Bitmap_cubical_complex.h +0 -0
  107. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Bitmap_cubical_complex_base.h +0 -0
  108. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Bitmap_cubical_complex_periodic_boundary_conditions_base.h +0 -0
  109. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Debug_utils.h +0 -0
  110. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Degree_rips_bifiltration.h +0 -0
  111. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Dynamic_multi_parameter_filtration.h +0 -0
  112. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Multi_field.h +0 -0
  113. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Multi_field_operators.h +0 -0
  114. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Multi_field_shared.h +0 -0
  115. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Multi_field_small.h +0 -0
  116. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Multi_field_small_operators.h +0 -0
  117. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Multi_field_small_shared.h +0 -0
  118. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Z2_field.h +0 -0
  119. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Z2_field_operators.h +0 -0
  120. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Zp_field.h +0 -0
  121. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Zp_field_operators.h +0 -0
  122. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Fields/Zp_field_shared.h +0 -0
  123. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Flag_complex_edge_collapser.h +0 -0
  124. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Matrix.h +0 -0
  125. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_filtration/Multi_parameter_generator.h +0 -0
  126. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_filtration/multi_filtration_conversions.h +0 -0
  127. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_filtration/multi_filtration_utils.h +0 -0
  128. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_parameter_filtered_complex.h +0 -0
  129. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_parameter_filtration.h +0 -0
  130. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_persistence/Box.h +0 -0
  131. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_persistence/Line.h +0 -0
  132. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_persistence/Multi_parameter_filtered_complex_pcoh_interface.h +0 -0
  133. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_persistence/Persistence_interface_cohomology.h +0 -0
  134. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_persistence/Persistence_interface_homology.h +0 -0
  135. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_persistence/Persistence_interface_vineyard.h +0 -0
  136. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Multi_persistence/Point.h +0 -0
  137. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Off_reader.h +0 -0
  138. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/Base_matrix.h +0 -0
  139. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/Base_matrix_with_column_compression.h +0 -0
  140. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/Boundary_matrix.h +0 -0
  141. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/Chain_matrix.h +0 -0
  142. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/Id_to_index_overlay.h +0 -0
  143. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/Position_to_index_overlay.h +0 -0
  144. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/RU_matrix.h +0 -0
  145. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/allocators/entry_constructors.h +0 -0
  146. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/base_pairing.h +0 -0
  147. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/base_swap.h +0 -0
  148. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/chain_pairing.h +0 -0
  149. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/chain_rep_cycles.h +0 -0
  150. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/chain_vine_swap.h +0 -0
  151. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/chain_column_extra_properties.h +0 -0
  152. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/column_dimension_holder.h +0 -0
  153. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/column_utilities.h +0 -0
  154. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/entry_types.h +0 -0
  155. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/heap_column.h +0 -0
  156. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/intrusive_list_column.h +0 -0
  157. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/intrusive_set_column.h +0 -0
  158. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/list_column.h +0 -0
  159. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/naive_vector_column.h +0 -0
  160. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/row_access.h +0 -0
  161. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/set_column.h +0 -0
  162. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/unordered_set_column.h +0 -0
  163. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/columns/vector_column.h +0 -0
  164. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/index_mapper.h +0 -0
  165. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/matrix_dimension_holders.h +0 -0
  166. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/matrix_row_access.h +0 -0
  167. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/ru_pairing.h +0 -0
  168. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/ru_rep_cycles.h +0 -0
  169. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_matrix/ru_vine_swap.h +0 -0
  170. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_on_a_line.h +0 -0
  171. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistence_on_rectangle.h +0 -0
  172. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistent_cohomology/Field_Zp.h +0 -0
  173. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistent_cohomology/Multi_field.h +0 -0
  174. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistent_cohomology/Persistent_cohomology_column.h +0 -0
  175. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Persistent_cohomology.h +0 -0
  176. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Points_off_io.h +0 -0
  177. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Projective_cover_kernel.h +0 -0
  178. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simple_object_pool.h +0 -0
  179. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simplex_tree/Simplex_tree_iterators.h +0 -0
  180. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simplex_tree/Simplex_tree_node_explicit_storage.h +0 -0
  181. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simplex_tree/Simplex_tree_siblings.h +0 -0
  182. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simplex_tree/Simplex_tree_star_simplex_iterators.h +0 -0
  183. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simplex_tree/filtration_value_utils.h +0 -0
  184. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simplex_tree/hooks_simplex_base.h +0 -0
  185. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simplex_tree/indexing_tag.h +0 -0
  186. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simplex_tree/serialization_utils.h +0 -0
  187. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simplex_tree/simplex_tree_options.h +0 -0
  188. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Simplex_tree.h +0 -0
  189. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/Thread_safe_slicer.h +0 -0
  190. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/distance_functions.h +0 -0
  191. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/graph_simplicial_complex.h +0 -0
  192. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/multi_simplex_tree_helpers.h +0 -0
  193. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/persistence_interval.h +0 -0
  194. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/persistence_matrix_options.h +0 -0
  195. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/reader_utils.h +0 -0
  196. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/simple_mdspan.h +0 -0
  197. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/slicer_helpers.h +0 -0
  198. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/gudhi/vineyard_base.h +0 -0
  199. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/tmp_h0_pers/mma_interface_h0.h +0 -0
  200. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/gudhi/tmp_h0_pers/naive_merge_tree.h +0 -0
  201. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/ml/__init__.py +0 -0
  202. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/ml/accuracies.py +0 -0
  203. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/ml/invariants_with_persistable.py +0 -0
  204. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/ml/kernels.py +0 -0
  205. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/ml/one.py +0 -0
  206. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/ml/tools.py +0 -0
  207. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multi_parameter_rank_invariant/diff_helpers.h +0 -0
  208. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multi_parameter_rank_invariant/euler_characteristic.h +0 -0
  209. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multi_parameter_rank_invariant/function_rips.h +0 -0
  210. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multi_parameter_rank_invariant/hilbert_function.h +0 -0
  211. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multi_parameter_rank_invariant/persistence_slices.h +0 -0
  212. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multi_parameter_rank_invariant/rank_invariant.h +0 -0
  213. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multiparameter_edge_collapse.py +0 -0
  214. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multiparameter_module_approximation/debug.h +0 -0
  215. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multiparameter_module_approximation/format_python-cpp.h +0 -0
  216. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/multiparameter_module_approximation/utilities.h +0 -0
  217. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/pickle.py +0 -0
  218. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/point_measure.pyx +0 -0
  219. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/simplex_tree_multi.pxd +0 -0
  220. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/tensor/tensor.h +0 -0
  221. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/tensor.pxd +0 -0
  222. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/test.pyx +0 -0
  223. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/tests/__init__.py +0 -0
  224. {multipers-2.4.1 → multipers-2.4.2b1}/multipers/vector_interface.pxd +0 -0
  225. {multipers-2.4.1 → multipers-2.4.2b1}/multipers.egg-info/dependency_links.txt +0 -0
  226. {multipers-2.4.1 → multipers-2.4.2b1}/multipers.egg-info/requires.txt +0 -0
  227. {multipers-2.4.1 → multipers-2.4.2b1}/multipers.egg-info/top_level.txt +0 -0
  228. {multipers-2.4.1 → multipers-2.4.2b1}/setup.cfg +0 -0
  229. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_aida.py +0 -0
  230. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_diff_helper.py +0 -0
  231. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_filtrations.py +0 -0
  232. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_hilbert_function.py +0 -0
  233. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_python-cpp_conversion.py +0 -0
  234. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_signed_betti.py +0 -0
  235. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_signed_measure.py +0 -0
  236. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_simplextreemulti.py +0 -0
  237. {multipers-2.4.1 → multipers-2.4.2b1}/tests/test_slicer.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: multipers
3
- Version: 2.4.1
3
+ Version: 2.4.2b1
4
4
  Summary: Multiparameter Topological Persistence for Machine Learning
5
5
  Author-email: David Loiseaux <david.lapous@proton.me>, Hannah Schreiber <hannah.schreiber@inria.fr>
6
6
  Maintainer-email: David Loiseaux <david.lapous@proton.me>
@@ -162,7 +162,7 @@ def signed_measure(
162
162
  grid = filtered_complex.filtration_grid
163
163
 
164
164
  if mass_default is None:
165
- mass_default = mass_default
165
+ pass
166
166
  elif isinstance(mass_default, str):
167
167
  if mass_default == "auto":
168
168
  mass_default = np.array([1.1 * np.max(f) - 0.1 * np.min(f) for f in grid])
@@ -228,7 +228,7 @@ def signed_measure(
228
228
  zero_pad=fix_mass_default,
229
229
  # grid_shape=tuple(len(g) for g in grid),
230
230
  ignore_inf=ignore_infinite_filtration_values,
231
- verbose=verbose
231
+ verbose=verbose,
232
232
  )
233
233
  fix_mass_default = False
234
234
  if verbose:
@@ -52,8 +52,7 @@ def _blocks2boundary_dimension_grades(
52
52
  b[0] if len(b[0]) > 0 else np.empty((0, num_parameters))
53
53
  for b in rblocks
54
54
  ),
55
- dtype=filtration_type,
56
- )
55
+ ).astype(filtration_type)
57
56
  boundary = tuple(x + S[i] for i, b in enumerate(rblocks) for x in b[1])
58
57
  dimensions = np.fromiter(
59
58
  (i for i, b in enumerate(rblocks) for _ in range(len(b[0]))), dtype=int
@@ -62,6 +61,7 @@ def _blocks2boundary_dimension_grades(
62
61
 
63
62
 
64
63
  def _slicer_from_simplextree(st, backend, vineyard):
64
+ backend = backend.lower() if isinstance(backend, str) else backend
65
65
  if vineyard:
66
66
  if backend == "matrix":
67
67
  slicer = mps._SlicerVineSimplicial(st)
@@ -172,7 +172,7 @@ def Slicer(
172
172
  else:
173
173
  vineyard = False if vineyard is None else vineyard
174
174
  column_type = "INTRUSIVE_SET" if column_type is None else column_type
175
- backend = "Matrix" if backend is None else backend
175
+ backend = "matrix" if backend is None else backend
176
176
 
177
177
  _Slicer = mps.get_matrix_slicer(
178
178
  is_vineyard=vineyard,
@@ -189,7 +189,7 @@ def Slicer(
189
189
  return _Slicer()
190
190
  elif mps.is_slicer(st):
191
191
  slicer = _Slicer(st)
192
- elif is_simplextree_multi(st) and backend == "Graph":
192
+ elif is_simplextree_multi(st) and backend == "graph":
193
193
  slicer = _slicer_from_simplextree(st, backend, vineyard)
194
194
  if st.is_squeezed:
195
195
  slicer.filtration_grid = st.filtration_grid
@@ -219,6 +219,7 @@ You can try using `multipers.slicer.to_simplextree`."""
219
219
  )
220
220
  if reduce:
221
221
  from multipers.ops import minimal_presentation
222
+
222
223
  slicer = minimal_presentation(
223
224
  slicer,
224
225
  backend=reduce_backend,
@@ -1,5 +1,5 @@
1
1
  import multipers.array_api.numpy as npapi
2
-
2
+ available_api = [npapi]
3
3
 
4
4
  def api_from_tensor(x, *, verbose: bool = False, strict=False):
5
5
  if strict:
@@ -24,8 +24,19 @@ searchsorted = _np.searchsorted
24
24
  LazyTensor = None
25
25
  abs = _np.abs
26
26
  exp = _np.exp
27
+ log = _np.log
27
28
  sin = _np.sin
28
29
  cos = _np.cos
30
+ matmul = _np.matmul
31
+ einsum = _np.einsum
32
+
33
+
34
+ def argsort(x, axis=-1):
35
+ return _np.argsort(x, axis=axis)
36
+
37
+
38
+ def astype(x, dtype):
39
+ return astensor(x).astype(dtype=dtype)
29
40
 
30
41
 
31
42
  def clip(x, min=None, max=None):
@@ -122,3 +133,26 @@ def is_promotable(x):
122
133
 
123
134
  def has_grad(_):
124
135
  return False
136
+
137
+
138
+ def to_device(x, device):
139
+ if device is None or str(device) in {"None", "cpu"}:
140
+ return x
141
+ raise ValueError(
142
+ f"NumPy backend only supports CPU tensors, requested device {device!r}."
143
+ )
144
+
145
+
146
+ def size(x):
147
+ return int(_np.size(x))
148
+
149
+
150
+ def dtype_is_float(dtype):
151
+ try:
152
+ return _np.issubdtype(_np.dtype(dtype), _np.floating)
153
+ except TypeError:
154
+ return False
155
+
156
+
157
+ def dtype_default():
158
+ return _np.array(0.0).dtype
@@ -1,5 +1,9 @@
1
1
  import numpy as _np
2
2
  import torch as _t
3
+ import multipers.array_api as _mpapi
4
+ import sys
5
+
6
+ _mpapi.available_api.append(sys.modules[__name__])
3
7
 
4
8
  backend = _t
5
9
  cat = _t.cat
@@ -23,8 +27,19 @@ LazyTensor = None
23
27
  relu = _t.relu
24
28
  abs = _t.abs
25
29
  exp = _t.exp
30
+ log = _t.log
26
31
  sin = _t.sin
27
32
  cos = _t.cos
33
+ matmul = _t.matmul
34
+ einsum = _t.einsum
35
+
36
+
37
+ def argsort(x, axis=-1):
38
+ return _t.argsort(x, dim=axis)
39
+
40
+
41
+ def astype(x, dtype):
42
+ return astensor(x).type(dtype)
28
43
 
29
44
 
30
45
  _is_keops_available = None
@@ -34,7 +49,6 @@ def clip(x, min=None, max=None):
34
49
  return _t.clamp(x, min, max)
35
50
 
36
51
 
37
-
38
52
  def split_with_sizes(arr, sizes):
39
53
  return arr.split_with_sizes(sizes)
40
54
 
@@ -131,3 +145,21 @@ def is_promotable(x):
131
145
 
132
146
  def has_grad(x):
133
147
  return x.requires_grad
148
+
149
+
150
+ def to_device(x, device):
151
+ if device is None:
152
+ return x
153
+ return x.to(device)
154
+
155
+
156
+ def size(x):
157
+ return x.numel()
158
+
159
+
160
+ def dtype_is_float(dtype):
161
+ return getattr(dtype, "is_floating_point", False)
162
+
163
+
164
+ def dtype_default():
165
+ return _t.get_default_dtype()
@@ -2,8 +2,8 @@ from collections.abc import Callable, Iterable
2
2
  from typing import Any, Literal, Union
3
3
 
4
4
  import numpy as np
5
-
6
- from multipers.array_api import api_from_tensor, api_from_tensors
5
+ import multipers.array_api.numpy as _npapi
6
+ from multipers.array_api import api_from_tensor
7
7
 
8
8
  global available_kernels
9
9
  available_kernels = Union[
@@ -30,7 +30,7 @@ def convolution_signed_measures(
30
30
  Parameters
31
31
  ----------
32
32
 
33
- - iterable_of_signed_measures : (num_signed_measure) x [ (npts) x (num_parameters), (npts)]
33
+ - iterable_of_signed_measures : (num_signed_measure) x [ (npts,num_parameters), (npts)]
34
34
  - filtrations : (num_parameter) x (filtration values)
35
35
  - flatten : bool
36
36
  - n_jobs : int
@@ -42,10 +42,14 @@ def convolution_signed_measures(
42
42
  """
43
43
  from multipers.grids import todense
44
44
 
45
- grid_iterator = todense(filtrations, product_order=True)
45
+ grid_iterator = todense(filtrations)
46
46
  api = api_from_tensor(iterable_of_signed_measures[0][0][0])
47
47
  match backend:
48
48
  case "sklearn":
49
+ if api is not _npapi:
50
+ raise ValueError(
51
+ f"The sklearn backend only supports numpy. Got {api=}."
52
+ )
49
53
 
50
54
  def convolution_signed_measures_on_grid(
51
55
  signed_measures,
@@ -78,6 +82,7 @@ def convolution_signed_measures(
78
82
  grid_iterator=grid_iterator,
79
83
  bandwidth=bandwidth,
80
84
  kernel=kernel,
85
+ api=api,
81
86
  **kwargs,
82
87
  )
83
88
  for pts, weights in signed_measures
@@ -95,6 +100,7 @@ def convolution_signed_measures(
95
100
  grid_iterator=grid_iterator,
96
101
  bandwidth=bandwidth,
97
102
  kernel=kernel,
103
+ api=api,
98
104
  **kwargs,
99
105
  )
100
106
 
@@ -143,8 +149,12 @@ def _pts_convolution_sparse_old(
143
149
  if len(pts) == 0:
144
150
  # warn("Found a trivial signed measure !")
145
151
  return np.zeros(len(grid_iterator))
152
+ if kernel == "multivariate_gaussian":
153
+ kernel = "gaussian"
154
+ if kernel == "sinc":
155
+ raise ValueError("Sinc kernel is not supported by sklearn backend.")
146
156
  kde = KernelDensity(
147
- kernel=kernel, bandwidth=bandwidth, rtol=1e-4, **more_kde_args
157
+ kernel=kernel, bandwidth=bandwidth, **more_kde_args
148
158
  ) # TODO : check rtol
149
159
  pos_indices = pts_weights > 0
150
160
  neg_indices = pts_weights < 0
@@ -166,31 +176,36 @@ def _pts_convolution_sparse_old(
166
176
 
167
177
 
168
178
  def _pts_convolution_pykeops(
169
- pts: np.ndarray,
170
- pts_weights: np.ndarray,
179
+ pts,
180
+ pts_weights,
171
181
  grid_iterator,
172
182
  kernel: available_kernels = "gaussian",
173
183
  bandwidth=0.1,
184
+ api=None,
174
185
  **more_kde_args,
175
186
  ):
176
187
  """
177
188
  Pykeops convolution
178
189
  """
179
- if isinstance(pts, np.ndarray):
180
- _asarray_weights = lambda x: np.asarray(x, dtype=pts.dtype)
181
- _asarray_grid = _asarray_weights
182
- else:
183
- import torch
184
-
185
- _asarray_weights = lambda x: torch.from_numpy(x).type(pts.dtype)
186
- _asarray_grid = lambda x: x.type(pts.dtype)
190
+ if api is None:
191
+ api = api_from_tensor(pts)
192
+ # if isinstance(pts, np.ndarray):
193
+ # _asarray_weights = lambda x: np.asarray(x, dtype=pts.dtype)
194
+ # _asarray_grid = _asarray_weights
195
+ # else:
196
+ # import torch
197
+ #
198
+ # _asarray_weights = lambda x: torch.from_numpy(x).type(pts.dtype)
199
+ # _asarray_grid = lambda x: x.type(pts.dtype)
200
+ pts = api.astensor(pts)
201
+ dtype=pts.dtype
187
202
  kde = KDE(kernel=kernel, bandwidth=bandwidth, **more_kde_args)
188
- return kde.fit(pts, sample_weights=_asarray_weights(pts_weights)).score_samples(
189
- _asarray_grid(grid_iterator)
203
+ return kde.fit(pts, sample_weights=api.astype(pts_weights, dtype),api=api).score_samples(
204
+ api.astype(grid_iterator, dtype)
190
205
  )
191
206
 
192
207
 
193
- def gaussian_kernel(x_i, y_j, bandwidth):
208
+ def gaussian_kernel(x_i, y_j, bandwidth, **kwargs):
194
209
  D = x_i.shape[-1]
195
210
  exponent = -(((x_i - y_j) / bandwidth) ** 2).sum(dim=-1) / 2
196
211
  # float is necessary for some reason (pykeops fails)
@@ -207,14 +222,14 @@ def multivariate_gaussian_kernel(x_i, y_j, covariance_matrix_inverse):
207
222
  exponent = -(z.weightedsqnorm(covariance_matrix_inverse.flatten()) / 2)
208
223
  return (
209
224
  float((2 * np.pi) ** (-dim / 2))
210
- * (covariance_matrix_inverse.det().sqrt())
225
+ * (np.sqrt(np.linalg.det(covariance_matrix_inverse)))
211
226
  * exponent.exp()
212
227
  )
213
228
 
214
229
 
215
- def exponential_kernel(x_i, y_j, bandwidth):
230
+ def exponential_kernel(x_i, y_j, bandwidth, **kwargs):
216
231
  # 1 / \sigma * exp( norm(x-y, dim=-1))
217
- exponent = -((((x_i - y_j) ** 2).sum(dim=-1) ** 1 / 2) / bandwidth)
232
+ exponent = -((((x_i - y_j) ** 2).sum(dim=-1) ** 0.5) / bandwidth)
218
233
  kernel = exponent.exp() / bandwidth
219
234
  return kernel
220
235
 
@@ -263,6 +278,7 @@ class KDE:
263
278
  bandwidth: Any = 1,
264
279
  kernel: available_kernels = "gaussian",
265
280
  return_log: bool = False,
281
+ **kwargs,
266
282
  ):
267
283
  """
268
284
  bandwidth : numeric
@@ -272,22 +288,15 @@ class KDE:
272
288
  self.bandwidth = bandwidth
273
289
  self.kernel: available_kernels = kernel
274
290
  self._kernel = None
275
- self._backend = None
291
+ self.api = None
276
292
  self._sample_weights = None
277
293
  self.return_log = return_log
294
+ self.kwargs = kwargs
278
295
 
279
- def fit(self, X, sample_weights=None, y=None):
296
+ def fit(self, X, sample_weights=None, y=None, api=None):
280
297
  self.X = X
281
298
  self._sample_weights = sample_weights
282
- if isinstance(X, np.ndarray):
283
- self._backend = np
284
- else:
285
- import torch
286
-
287
- if isinstance(X, torch.Tensor):
288
- self._backend = torch
289
- else:
290
- raise Exception("Unsupported backend.")
299
+ self.api = api_from_tensor(X) if api is None else api
291
300
  self._kernel = _kernel(self.kernel)
292
301
  return self
293
302
 
@@ -340,22 +349,20 @@ class KDE:
340
349
  log_probs : tensor (m)
341
350
  log probability densities for each of the queried points in `Y`
342
351
  """
343
- assert self._backend is not None and self._kernel is not None, "Fit first."
352
+ assert self.api is not None and self._kernel is not None, "Fit first."
344
353
  X = self.X if X is None else X
345
354
  if X.shape[0] == 0:
346
- return self._backend.zeros((Y.shape[0]))
355
+ return self.api.zeros((Y.shape[0]))
347
356
  assert Y.shape[1] == X.shape[1] and X.ndim == Y.ndim == 2
348
357
  lazy_x, lazy_y, w = self.to_lazy(X, Y, x_weights=self._sample_weights)
349
- kernel = self._kernel(lazy_x, lazy_y, self.bandwidth)
358
+ kernel = self._kernel(lazy_x, lazy_y, self.bandwidth, **self.kwargs)
350
359
  if w is not None:
351
360
  kernel *= w
352
361
  if return_kernel:
353
362
  return kernel
354
363
  density_estimation = kernel.sum(dim=0).squeeze() / kernel.shape[0] # mean
355
364
  return (
356
- self._backend.log(density_estimation)
357
- if self.return_log
358
- else density_estimation
365
+ self.api.log(density_estimation) if self.return_log else density_estimation
359
366
  )
360
367
 
361
368
 
@@ -377,7 +384,7 @@ class DTM:
377
384
  self._ks = None
378
385
  self._kdtree = None
379
386
  self._X = None
380
- self._backend = None
387
+ self.api = None
381
388
 
382
389
  def fit(self, X, sample_weights=None, y=None):
383
390
  if len(self.masses) == 0:
@@ -385,15 +392,8 @@ class DTM:
385
392
  assert np.max(self.masses) <= 1, "All masses should be in (0,1]."
386
393
  from sklearn.neighbors import KDTree
387
394
 
388
- if not isinstance(X, np.ndarray):
389
- import torch
390
-
391
- assert isinstance(X, torch.Tensor), "Backend has to be numpy of torch"
392
- _X = X.detach()
393
- self._backend = "torch"
394
- else:
395
- _X = X
396
- self._backend = "numpy"
395
+ self.api = api_from_tensor(X)
396
+ _X = self.api.asnumpy(X)
397
397
  self._ks = np.array([int(mass * X.shape[0]) + 1 for mass in self.masses])
398
398
  self._kdtree = KDTree(_X, metric=self.metric, **self._kdtree_kwargs)
399
399
  self._X = X
@@ -413,19 +413,21 @@ class DTM:
413
413
  -------
414
414
  the DTMs of Y, for each mass in masses.
415
415
  """
416
+ if self.api is None:
417
+ raise ValueError("Fit first")
416
418
  if len(self.masses) == 0:
417
- return np.empty((0, len(Y)))
419
+ return self.api.empty((0, len(Y)))
418
420
  assert (
419
421
  self._ks is not None and self._kdtree is not None and self._X is not None
420
422
  ), f"Fit first. Got {self._ks=}, {self._kdtree=}, {self._X=}."
421
423
  assert Y.ndim == 2
422
- if self._backend == "torch":
423
- _Y = Y.detach().numpy()
424
- else:
424
+ if self.api is _npapi:
425
425
  _Y = Y
426
+ else:
427
+ _Y = self.api.asnumpy(Y)
426
428
  NN_Dist, NN = self._kdtree.query(_Y, self._ks.max(), return_distance=True)
427
429
  DTMs = np.array([((NN_Dist**2)[:, :k].mean(1)) ** 0.5 for k in self._ks])
428
- return DTMs
430
+ return self.api.astensor(DTMs)
429
431
 
430
432
  def score_samples_diff(self, Y):
431
433
  """Returns the kernel density estimates of each point in `Y`.
@@ -447,12 +449,13 @@ class DTM:
447
449
  log probability densities for each of the queried points in `Y`
448
450
  """
449
451
  import torch
452
+ from multipers.array_api import torch
450
453
 
451
454
  if len(self.masses) == 0:
452
455
  return torch.empty(0, len(Y))
453
456
 
454
457
  assert Y.ndim == 2
455
- assert self._backend == "torch", "Use the non-diff version with numpy."
458
+ assert self.api is torch, "Use the non-diff version with numpy."
456
459
  assert (
457
460
  self._ks is not None and self._kdtree is not None and self._X is not None
458
461
  ), f"Fit first. Got {self._ks=}, {self._kdtree=}, {self._X=}."
@@ -43,17 +43,6 @@ def get_degree_rips(st, vector[int] degrees):
43
43
  with nogil:
44
44
  get_degree_rips_st_python(simplextree_ptr, st_multi_ptr, degrees)
45
45
  return degree_rips_st
46
- # cdef int max_degree = out.second
47
- # cdef bool inf_flag = filtrations[-1] == np.inf
48
- # if inf_flag:
49
- # filtrations = filtrations[:-1]
50
- # filtrations, = mpg.compute_grid([filtrations],strategy=grid_strategy,resolution=resolution)
51
- # if inf_flag:
52
- # filtrations = np.concatenate([filtrations, [np.inf]])
53
- # degree_rips_st.grid_squeeze([filtrations, degrees], inplace=True, coordinate_values=True)
54
- # degree_rips_st.filtration_grid = mpg.sanitize_grid([filtrations, np.asarray(degrees)])
55
- # # degree_rips_st._is_function_simplextree=True
56
- # return degree_rips_st,max_degree
57
46
 
58
47
  def function_rips_surface(st_multi, vector[indices_type] homological_degrees, bool mobius_inversion=True, bool zero_pad=False, indices_type n_jobs=0):
59
48
  assert st_multi.is_squeezed, "Squeeze first !"