TACE 0.2.0__tar.gz → 0.2.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 (369) hide show
  1. {tace-0.2.0 → tace-0.2.2}/PKG-INFO +14 -3
  2. tace-0.2.2/README.md +217 -0
  3. {tace-0.2.0 → tace-0.2.2}/TACE.egg-info/PKG-INFO +14 -3
  4. tace-0.2.2/TACE.egg-info/SOURCES.txt +308 -0
  5. {tace-0.2.0 → tace-0.2.2}/TACE.egg-info/entry_points.txt +3 -1
  6. {tace-0.2.0 → tace-0.2.2}/TACE.egg-info/requires.txt +17 -2
  7. {tace-0.2.0 → tace-0.2.2}/TACE.egg-info/top_level.txt +1 -0
  8. tace-0.2.2/eqx/LICENSE.md +402 -0
  9. tace-0.2.2/eqx/README.md +116 -0
  10. tace-0.2.2/eqx/__init__.py +31 -0
  11. tace-0.2.2/eqx/_layout.py +30 -0
  12. tace-0.2.2/eqx/ace/__init__.py +5 -0
  13. tace-0.2.2/eqx/ace/contraction.py +82 -0
  14. tace-0.2.2/eqx/ace/cuda.py +251 -0
  15. tace-0.2.2/eqx/ace/tace.py +273 -0
  16. tace-0.2.2/eqx/co2/__init__.py +43 -0
  17. tace-0.2.2/eqx/co2/basis.py +125 -0
  18. tace-0.2.2/eqx/co2/cartesian_harmonics.py +68 -0
  19. tace-0.2.2/eqx/co2/gate.py +72 -0
  20. tace-0.2.2/eqx/co2/irreps.py +247 -0
  21. tace-0.2.2/eqx/co2/linear.py +173 -0
  22. tace-0.2.2/eqx/co2/o3_tensor_product.py +324 -0
  23. tace-0.2.2/eqx/co2/restriction.py +265 -0
  24. tace-0.2.2/eqx/co2/spherical.py +256 -0
  25. tace-0.2.2/eqx/co2/tensor_product.py +282 -0
  26. tace-0.2.2/eqx/co3/__init__.py +28 -0
  27. tace-0.2.2/eqx/co3/basis.py +155 -0
  28. tace-0.2.2/eqx/co3/cartesian_harmonics.py +123 -0
  29. tace-0.2.2/eqx/co3/gate.py +104 -0
  30. tace-0.2.2/eqx/co3/irreps.py +288 -0
  31. tace-0.2.2/eqx/co3/linear.py +268 -0
  32. tace-0.2.2/eqx/co3/symmetric.py +207 -0
  33. tace-0.2.2/eqx/co3/tensor_product.py +422 -0
  34. tace-0.2.2/eqx/conv/__init__.py +18 -0
  35. tace-0.2.2/eqx/conv/angular.py +101 -0
  36. tace-0.2.2/eqx/conv/attention.py +477 -0
  37. tace-0.2.2/eqx/conv/co3/__init__.py +6 -0
  38. tace-0.2.2/eqx/conv/co3/convolution.py +431 -0
  39. tace-0.2.2/eqx/conv/co3/linear.py +213 -0
  40. tace-0.2.2/eqx/conv/co3/polynomials.py +365 -0
  41. tace-0.2.2/eqx/conv/codegen.py +455 -0
  42. tace-0.2.2/eqx/conv/contraction.py +47 -0
  43. tace-0.2.2/eqx/conv/edge.py +181 -0
  44. tace-0.2.2/eqx/conv/graph.py +94 -0
  45. tace-0.2.2/eqx/conv/o2_o3/__init__.py +5 -0
  46. tace-0.2.2/eqx/conv/o2_o3/autotune.py +140 -0
  47. tace-0.2.2/eqx/conv/o2_o3/codegen.py +570 -0
  48. tace-0.2.2/eqx/conv/o2_o3/convolution.py +598 -0
  49. tace-0.2.2/eqx/conv/o2_o3/cuda.py +630 -0
  50. tace-0.2.2/eqx/conv/o2_o3/direction_codegen.py +433 -0
  51. tace-0.2.2/eqx/conv/o2_o3/geometry.py +369 -0
  52. tace-0.2.2/eqx/conv/o2_o3/schedule.py +209 -0
  53. tace-0.2.2/eqx/conv/o2_o3/transverse.py +512 -0
  54. tace-0.2.2/eqx/conv/o3/__init__.py +5 -0
  55. tace-0.2.2/eqx/conv/o3/codegen.py +455 -0
  56. tace-0.2.2/eqx/conv/o3/convolution.py +383 -0
  57. tace-0.2.2/eqx/conv/o3/cuda.py +346 -0
  58. tace-0.2.2/eqx/conv/o3/harmonics.py +550 -0
  59. tace-0.2.2/eqx/conv/program.py +244 -0
  60. tace-0.2.2/eqx/conv/radial.py +183 -0
  61. tace-0.2.2/eqx/conv/uu_o2/__init__.py +5 -0
  62. tace-0.2.2/eqx/conv/uu_o2/convolution.py +385 -0
  63. tace-0.2.2/eqx/conv/uv_o2/__init__.py +5 -0
  64. tace-0.2.2/eqx/conv/uv_o2/convolution.py +645 -0
  65. tace-0.2.2/eqx/conv/uv_o2/transverse.py +176 -0
  66. tace-0.2.2/eqx/kernels/__init__.py +5 -0
  67. tace-0.2.2/eqx/kernels/channel_product.py +150 -0
  68. tace-0.2.2/eqx/kernels/codegen.py +401 -0
  69. tace-0.2.2/eqx/kernels/csrc/runtime.cpp +174 -0
  70. tace-0.2.2/eqx/kernels/cuda.py +181 -0
  71. tace-0.2.2/eqx/kernels/cuda_graph.py +170 -0
  72. tace-0.2.2/eqx/kernels/layout.py +226 -0
  73. tace-0.2.2/eqx/kernels/quaternion.py +347 -0
  74. tace-0.2.2/eqx/kernels/recompute.py +180 -0
  75. tace-0.2.2/eqx/kernels/rotary.py +161 -0
  76. tace-0.2.2/eqx/kernels/rotation.py +135 -0
  77. tace-0.2.2/eqx/kernels/wigner.py +277 -0
  78. tace-0.2.2/eqx/models/__init__.py +1 -0
  79. tace-0.2.2/eqx/models/convolution.py +274 -0
  80. tace-0.2.2/eqx/models/equflash/__init__.py +5 -0
  81. tace-0.2.2/eqx/models/equflash/conversion.py +198 -0
  82. tace-0.2.2/eqx/models/mace/__init__.py +5 -0
  83. tace-0.2.2/eqx/models/mace/conversion.py +217 -0
  84. tace-0.2.2/eqx/models/nequip/__init__.py +5 -0
  85. tace-0.2.2/eqx/models/nequip/conversion.py +145 -0
  86. tace-0.2.2/eqx/models/prophet/__init__.py +5 -0
  87. tace-0.2.2/eqx/models/prophet/conversion.py +110 -0
  88. tace-0.2.2/eqx/models/sevennet/__init__.py +5 -0
  89. tace-0.2.2/eqx/models/sevennet/conversion.py +136 -0
  90. tace-0.2.2/eqx/models/tace/__init__.py +1 -0
  91. tace-0.2.2/eqx/models/tace/tece_oam_rra/__init__.py +6 -0
  92. tace-0.2.2/eqx/models/tace/tece_oam_rra/bilinear_contraction.py +91 -0
  93. tace-0.2.2/eqx/models/tace/tece_oam_rra/bilinear_cuda.py +359 -0
  94. tace-0.2.2/eqx/models/tace/tece_oam_rra/cuda.py +34 -0
  95. tace-0.2.2/eqx/models/tace/tece_oam_rra/execution.py +265 -0
  96. tace-0.2.2/eqx/models/tace/tece_oam_rra/interaction.py +203 -0
  97. tace-0.2.2/eqx/models/tace/tece_oam_rra/product.py +248 -0
  98. tace-0.2.2/eqx/models/tace/tece_oam_rra/program.py +338 -0
  99. tace-0.2.2/eqx/o2/__init__.py +35 -0
  100. tace-0.2.2/eqx/o2/_clebsch_gordan.py +91 -0
  101. tace-0.2.2/eqx/o2/_layout.py +34 -0
  102. tace-0.2.2/eqx/o2/asymmetric_contraction.py +385 -0
  103. tace-0.2.2/eqx/o2/circular_harmonics.py +131 -0
  104. tace-0.2.2/eqx/o2/gate.py +434 -0
  105. tace-0.2.2/eqx/o2/irreps.py +625 -0
  106. tace-0.2.2/eqx/o2/linear.py +579 -0
  107. tace-0.2.2/eqx/o2/local_frame.py +421 -0
  108. tace-0.2.2/eqx/o2/o3_tensor_product.py +500 -0
  109. tace-0.2.2/eqx/o2/rotation_matrix.py +143 -0
  110. tace-0.2.2/eqx/o2/tensor_product.py +480 -0
  111. tace-0.2.2/eqx/o2/wigner.py +214 -0
  112. tace-0.2.2/eqx/o3/__init__.py +6 -0
  113. tace-0.2.2/eqx/o3/contraction.py +51 -0
  114. tace-0.2.2/eqx/o3/cuda.py +106 -0
  115. tace-0.2.2/eqx/o3/gate.py +108 -0
  116. tace-0.2.2/eqx/o3/linear.py +181 -0
  117. tace-0.2.2/eqx/utils.py +31 -0
  118. {tace-0.2.0 → tace-0.2.2}/pyproject.toml +63 -6
  119. {tace-0.2.0 → tace-0.2.2}/tace/__init__.py +4 -1
  120. tace-0.2.2/tace/dataset/augmentation.py +130 -0
  121. {tace-0.2.0 → tace-0.2.2}/tace/dataset/dataloader.py +88 -49
  122. {tace-0.2.0 → tace-0.2.2}/tace/dataset/datamodule.py +82 -66
  123. {tace-0.2.0 → tace-0.2.2}/tace/dataset/element.py +4 -5
  124. {tace-0.2.0 → tace-0.2.2}/tace/dataset/graph.py +13 -17
  125. {tace-0.2.0 → tace-0.2.2}/tace/dataset/neighbour_list.py +5 -5
  126. {tace-0.2.0 → tace-0.2.2}/tace/dataset/quantity.py +255 -211
  127. {tace-0.2.0 → tace-0.2.2}/tace/dataset/read.py +57 -36
  128. {tace-0.2.0 → tace-0.2.2}/tace/dataset/sampler.py +8 -4
  129. {tace-0.2.0 → tace-0.2.2}/tace/dataset/split.py +6 -7
  130. tace-0.2.2/tace/dataset/statistics.py +732 -0
  131. {tace-0.2.0 → tace-0.2.2}/tace/dataset/utils.py +19 -16
  132. {tace-0.2.0 → tace-0.2.2}/tace/foundations/__init__.py +1 -1
  133. {tace-0.2.0 → tace-0.2.2}/tace/foundations/download_link.py +9 -18
  134. {tace-0.2.0 → tace-0.2.2}/tace/foundations/u_shift.py +4 -3
  135. tace-0.2.2/tace/interface/ase/__init__.py +10 -0
  136. {tace-0.2.0 → tace-0.2.2}/tace/interface/ase/calculator.py +29 -31
  137. tace-0.2.2/tace/interface/ase/general_calculator.py +81 -0
  138. tace-0.2.2/tace/interface/ase/optimizer.py +540 -0
  139. tace-0.2.2/tace/interface/nvalchemi/__init__.py +4 -0
  140. tace-0.2.2/tace/interface/nvalchemi/wrapper.py +397 -0
  141. {tace-0.2.0 → tace-0.2.2}/tace/interface/torchsim/__init__.py +1 -1
  142. {tace-0.2.0 → tace-0.2.2}/tace/interface/torchsim/torchsim.py +85 -67
  143. {tace-0.2.0 → tace-0.2.2}/tace/interface/uspex26/relax.py +5 -4
  144. {tace-0.2.0 → tace-0.2.2}/tace/lightning/__init__.py +2 -1
  145. {tace-0.2.0 → tace-0.2.2}/tace/lightning/lit_model.py +276 -142
  146. {tace-0.2.0 → tace-0.2.2}/tace/lightning/lora.py +0 -1
  147. tace-0.2.0/tace/lightning/los_skip.py → tace-0.2.2/tace/lightning/loss_skip.py +6 -9
  148. tace-0.2.2/tace/lightning/torch_model.py +114 -0
  149. {tace-0.2.0 → tace-0.2.2}/tace/lightning/trainer.py +31 -24
  150. tace-0.2.2/tace/lightning/u_shift.py +141 -0
  151. tace-0.2.2/tace/models/__init__.py +21 -0
  152. {tace-0.2.0 → tace-0.2.2}/tace/models/_e3nn/__init__.py +1 -1
  153. {tace-0.2.0 → tace-0.2.2}/tace/models/_e3nn/base.py +171 -88
  154. tace-0.2.2/tace/models/_e3nn/basis_change.py +48 -0
  155. tace-0.2.2/tace/models/_e3nn/default.py +334 -0
  156. {tace-0.2.0 → tace-0.2.2}/tace/models/_e3nn/dropout.py +76 -80
  157. {tace-0.2.0 → tace-0.2.2}/tace/models/_e3nn/edge.py +22 -113
  158. tace-0.2.2/tace/models/_e3nn/fused.py +634 -0
  159. tace-0.2.2/tace/models/_e3nn/inter.py +895 -0
  160. tace-0.2.2/tace/models/_e3nn/layer_norm.py +171 -0
  161. tace-0.2.2/tace/models/_e3nn/les.py +395 -0
  162. tace-0.2.2/tace/models/_e3nn/magnetic.py +201 -0
  163. tace-0.2.2/tace/models/_e3nn/node.py +355 -0
  164. {tace-0.2.0 → tace-0.2.2}/tace/models/_e3nn/nonlinear.py +113 -117
  165. tace-0.2.2/tace/models/_e3nn/o2.py +683 -0
  166. tace-0.2.2/tace/models/_e3nn/paths.py +176 -0
  167. tace-0.2.2/tace/models/_e3nn/prod.py +498 -0
  168. tace-0.2.2/tace/models/_e3nn/readout.py +316 -0
  169. tace-0.2.2/tace/models/_e3nn/representation.py +679 -0
  170. {tace-0.2.0 → tace-0.2.2}/tace/models/_e3nn/residual.py +10 -11
  171. {tace-0.2.0 → tace-0.2.2}/tace/models/_e3nn/scatter_norm.py +6 -6
  172. tace-0.2.2/tace/models/_e3nn/symmetric_contraction.py +565 -0
  173. {tace-0.2.0 → tace-0.2.2}/tace/models/_e3nn/tace.py +322 -268
  174. tace-0.2.2/tace/models/_e3nn/tece_oam_rra.py +531 -0
  175. {tace-0.2.0 → tace-0.2.2}/tace/models/_e3nn/ue.py +51 -24
  176. tace-0.2.2/tace/models/_e3nn/wigner6j.py +553 -0
  177. {tace-0.2.0 → tace-0.2.2}/tace/models/adapter.py +108 -96
  178. {tace-0.2.0 → tace-0.2.2}/tace/models/angular.py +70 -71
  179. tace-0.2.2/tace/models/blocks.py +237 -0
  180. {tace-0.2.0 → tace-0.2.2}/tace/models/compile/aot.py +98 -48
  181. {tace-0.2.0 → tace-0.2.2}/tace/models/compile/compile.py +40 -6
  182. {tace-0.2.0 → tace-0.2.2}/tace/models/compile/tace.py +5 -0
  183. {tace-0.2.0 → tace-0.2.2}/tace/models/compile/wrapper.py +80 -23
  184. {tace-0.2.0/tace/models/_cue → tace-0.2.2/tace/models/cue}/__init__.py +2 -2
  185. {tace-0.2.0/tace/models/_cue → tace-0.2.2/tace/models/cue}/_tp_scatter.py +2 -4
  186. {tace-0.2.0/tace/models/_cue → tace-0.2.2/tace/models/cue}/paths.py +23 -15
  187. tace-0.2.2/tace/models/eqt/__init__.py +17 -0
  188. {tace-0.2.0/tace/models/_eqt → tace-0.2.2/tace/models/eqt}/_tp_uuu.py +15 -8
  189. {tace-0.2.0/tace/models/_eqt → tace-0.2.2/tace/models/eqt}/equitorch/irreps/__init__.py +1 -1
  190. {tace-0.2.0/tace/models/_eqt → tace-0.2.2/tace/models/eqt}/equitorch/nn/sparse_product.py +18 -12
  191. {tace-0.2.0/tace/models/_eqt → tace-0.2.2/tace/models/eqt}/equitorch/nn/tensor_products.py +36 -20
  192. {tace-0.2.0/tace/models/_eqt → tace-0.2.2/tace/models/eqt}/equitorch/structs/__init__.py +2 -1
  193. {tace-0.2.0 → tace-0.2.2}/tace/models/ictd.py +82 -83
  194. tace-0.2.2/tace/models/kspace.py +194 -0
  195. {tace-0.2.0 → tace-0.2.2}/tace/models/lammps.py +7 -10
  196. tace-0.2.2/tace/models/layout.py +138 -0
  197. {tace-0.2.0 → tace-0.2.2}/tace/models/linear.py +191 -82
  198. tace-0.2.2/tace/models/mlp.py +192 -0
  199. {tace-0.2.0 → tace-0.2.2}/tace/models/normalizer.py +3 -3
  200. tace-0.2.2/tace/models/oeq/__init__.py +6 -0
  201. {tace-0.2.0/tace/models/_oeq → tace-0.2.2/tace/models/oeq}/_tp_scatter.py +31 -24
  202. {tace-0.2.0 → tace-0.2.2}/tace/models/radial.py +404 -402
  203. tace-0.2.2/tace/models/s2.py +377 -0
  204. tace-0.2.2/tace/models/scf/README.md +116 -0
  205. tace-0.2.2/tace/models/scf/__init__.py +34 -0
  206. tace-0.2.2/tace/models/scf/electrostatics.py +211 -0
  207. tace-0.2.2/tace/models/scf/energy_functional.py +170 -0
  208. tace-0.2.2/tace/models/scf/fixed_point.py +299 -0
  209. tace-0.2.2/tace/models/scf/local.py +188 -0
  210. tace-0.2.2/tace/models/scf/longrange/LICENSE.md +11 -0
  211. tace-0.2.2/tace/models/scf/longrange/NOTICE.md +6 -0
  212. tace-0.2.2/tace/models/scf/longrange/__init__.py +5 -0
  213. tace-0.2.2/tace/models/scf/longrange/__version__.py +1 -0
  214. tace-0.2.2/tace/models/scf/longrange/energy.py +194 -0
  215. tace-0.2.2/tace/models/scf/longrange/features.py +1013 -0
  216. tace-0.2.2/tace/models/scf/longrange/gto_utils.py +413 -0
  217. tace-0.2.2/tace/models/scf/longrange/kspace.py +251 -0
  218. tace-0.2.2/tace/models/scf/longrange/realspace_electrostatics.py +422 -0
  219. tace-0.2.2/tace/models/scf/longrange/realspace_grid_integrals.py +202 -0
  220. tace-0.2.2/tace/models/scf/longrange/slabs.py +246 -0
  221. tace-0.2.2/tace/models/scf/longrange/utils.py +100 -0
  222. tace-0.2.2/tace/models/scf/loss.py +84 -0
  223. tace-0.2.2/tace/models/scf/model.py +74 -0
  224. tace-0.2.2/tace/models/scf/qeq.py +132 -0
  225. tace-0.2.2/tace/models/scf/readout.py +147 -0
  226. tace-0.2.2/tace/models/scf/state.py +71 -0
  227. tace-0.2.2/tace/models/scf/utils.py +115 -0
  228. {tace-0.2.0 → tace-0.2.2}/tace/models/softmax.py +119 -122
  229. tace-0.2.2/tace/models/time_reversal.py +60 -0
  230. {tace-0.2.0 → tace-0.2.2}/tace/models/utils.py +71 -25
  231. {tace-0.2.0 → tace-0.2.2}/tace/models/zemin.py +61 -61
  232. {tace-0.2.0 → tace-0.2.2}/tace/scripts/average.py +8 -10
  233. {tace-0.2.0 → tace-0.2.2}/tace/scripts/clean.py +1 -3
  234. tace-0.2.2/tace/scripts/convert_cgtp.py +41 -0
  235. tace-0.2.0/tace/scripts/convert.py → tace-0.2.2/tace/scripts/convert_lora.py +15 -11
  236. {tace-0.2.0 → tace-0.2.2}/tace/scripts/eval.py +88 -37
  237. {tace-0.2.0 → tace-0.2.2}/tace/scripts/export_eval.py +17 -3
  238. {tace-0.2.0 → tace-0.2.2}/tace/scripts/export_lammps.py +8 -4
  239. {tace-0.2.0 → tace-0.2.2}/tace/scripts/export_train.py +3 -1
  240. {tace-0.2.0 → tace-0.2.2}/tace/scripts/finetune.py +26 -17
  241. tace-0.2.2/tace/scripts/plot_diatom.py +296 -0
  242. {tace-0.2.0 → tace-0.2.2}/tace/scripts/split.py +1 -2
  243. {tace-0.2.0 → tace-0.2.2}/tace/scripts/train.py +72 -95
  244. {tace-0.2.0 → tace-0.2.2}/tace/scripts/update.py +21 -19
  245. tace-0.2.2/tace/scripts/utils/__init__.py +0 -0
  246. tace-0.2.2/tace/scripts/utils/check_equi.py +132 -0
  247. tace-0.2.2/tace/scripts/utils/check_soc.py +262 -0
  248. tace-0.2.2/tace/scripts/utils/convert_to_xyz.py +215 -0
  249. {tace-0.2.0 → tace-0.2.2}/tace/utils/__init__.py +2 -2
  250. {tace-0.2.0 → tace-0.2.2}/tace/utils/_global.py +4 -5
  251. {tace-0.2.0 → tace-0.2.2}/tace/utils/callbacks.py +5 -8
  252. {tace-0.2.0 → tace-0.2.2}/tace/utils/cfg.py +1 -3
  253. {tace-0.2.0 → tace-0.2.2}/tace/utils/ema.py +19 -19
  254. {tace-0.2.0 → tace-0.2.2}/tace/utils/env.py +41 -24
  255. {tace-0.2.0 → tace-0.2.2}/tace/utils/hydra_resolver.py +1 -2
  256. {tace-0.2.0 → tace-0.2.2}/tace/utils/logger.py +3 -8
  257. {tace-0.2.0 → tace-0.2.2}/tace/utils/loss/__init__.py +21 -22
  258. {tace-0.2.0 → tace-0.2.2}/tace/utils/loss/common.py +28 -1
  259. {tace-0.2.0 → tace-0.2.2}/tace/utils/loss/dens.py +562 -544
  260. {tace-0.2.0 → tace-0.2.2}/tace/utils/loss/huber_fn.py +571 -438
  261. {tace-0.2.0 → tace-0.2.2}/tace/utils/loss/l2mae_fn.py +119 -87
  262. {tace-0.2.0 → tace-0.2.2}/tace/utils/loss/mae_fn.py +161 -101
  263. {tace-0.2.0 → tace-0.2.2}/tace/utils/loss/mse_fn.py +161 -99
  264. {tace-0.2.0 → tace-0.2.2}/tace/utils/loss/normal.py +67 -76
  265. {tace-0.2.0 → tace-0.2.2}/tace/utils/loss/registry.py +41 -10
  266. {tace-0.2.0 → tace-0.2.2}/tace/utils/loss/special_fn.py +90 -91
  267. tace-0.2.2/tace/utils/loss/uncertainty.py +79 -0
  268. {tace-0.2.0 → tace-0.2.2}/tace/utils/lr_scheduler/__init__.py +1 -1
  269. {tace-0.2.0 → tace-0.2.2}/tace/utils/lr_scheduler/warmup.py +1 -2
  270. {tace-0.2.0 → tace-0.2.2}/tace/utils/lr_scheduler/wsd.py +33 -14
  271. {tace-0.2.0 → tace-0.2.2}/tace/utils/metrics.py +35 -28
  272. {tace-0.2.0 → tace-0.2.2}/tace/utils/metrics_bake.py +42 -28
  273. {tace-0.2.0 → tace-0.2.2}/tace/utils/optimizer/__init__.py +1 -1
  274. {tace-0.2.0 → tace-0.2.2}/tace/utils/optimizer/hybrid_muon.py +4 -4
  275. tace-0.2.2/tace/utils/optimizer/soap.py +1 -0
  276. {tace-0.2.0 → tace-0.2.2}/tace/utils/spectra.py +1 -2
  277. {tace-0.2.0 → tace-0.2.2}/tace/utils/strategy.py +4 -3
  278. {tace-0.2.0 → tace-0.2.2}/tace/utils/torch_scatter.py +55 -42
  279. {tace-0.2.0 → tace-0.2.2}/tace/utils/units.py +12 -12
  280. {tace-0.2.0 → tace-0.2.2}/tace/utils/utils.py +13 -71
  281. tace-0.2.2/tests/test_cartesian.py +121 -0
  282. tace-0.2.2/tests/test_compile.py +370 -0
  283. tace-0.2.2/tests/test_eqt.py +286 -0
  284. tace-0.2.2/tests/test_eqx_integration.py +1487 -0
  285. tace-0.2.2/tests/test_les.py +124 -0
  286. tace-0.2.2/tests/test_local_o2.py +1535 -0
  287. tace-0.2.2/tests/test_loss.py +142 -0
  288. tace-0.2.2/tests/test_magnetic_optimizer.py +96 -0
  289. tace-0.2.2/tests/test_product.py +462 -0
  290. tace-0.2.2/tests/test_radial.py +250 -0
  291. tace-0.2.2/tests/test_scf.py +167 -0
  292. tace-0.2.2/tests/test_statistics.py +416 -0
  293. tace-0.2.2/tests/test_time_reversal.py +613 -0
  294. tace-0.2.2/tests/test_torchsim.py +169 -0
  295. tace-0.2.2/tests/test_wigner6j.py +495 -0
  296. tace-0.2.0/README.rst +0 -171
  297. tace-0.2.0/TACE.egg-info/SOURCES.txt +0 -174
  298. tace-0.2.0/tace/dataset/statistics.py +0 -457
  299. tace-0.2.0/tace/interface/ase/__init__.py +0 -3
  300. tace-0.2.0/tace/lightning/torch_model.py +0 -107
  301. tace-0.2.0/tace/lightning/u_shift.py +0 -59
  302. tace-0.2.0/tace/models/__init__.py +0 -14
  303. tace-0.2.0/tace/models/_cart/__init__.py +0 -7
  304. tace-0.2.0/tace/models/_cart/base.py +0 -24
  305. tace-0.2.0/tace/models/_cart/basis_change.py +0 -118
  306. tace-0.2.0/tace/models/_cart/cartesian.py +0 -150
  307. tace-0.2.0/tace/models/_cart/default.py +0 -216
  308. tace-0.2.0/tace/models/_cart/dropout.py +0 -8
  309. tace-0.2.0/tace/models/_cart/edge.py +0 -26
  310. tace-0.2.0/tace/models/_cart/fused.py +0 -531
  311. tace-0.2.0/tace/models/_cart/inter.py +0 -201
  312. tace-0.2.0/tace/models/_cart/layer_norm.py +0 -108
  313. tace-0.2.0/tace/models/_cart/node.py +0 -105
  314. tace-0.2.0/tace/models/_cart/nonlinear.py +0 -8
  315. tace-0.2.0/tace/models/_cart/paths.py +0 -65
  316. tace-0.2.0/tace/models/_cart/prod.py +0 -238
  317. tace-0.2.0/tace/models/_cart/readout.py +0 -20
  318. tace-0.2.0/tace/models/_cart/representation.py +0 -354
  319. tace-0.2.0/tace/models/_cart/residual.py +0 -8
  320. tace-0.2.0/tace/models/_cart/tace.py +0 -550
  321. tace-0.2.0/tace/models/_cart/ue.py +0 -8
  322. tace-0.2.0/tace/models/_e3nn/asymmetric_contraction.py +0 -1140
  323. tace-0.2.0/tace/models/_e3nn/basis_change.py +0 -113
  324. tace-0.2.0/tace/models/_e3nn/default.py +0 -283
  325. tace-0.2.0/tace/models/_e3nn/fused.py +0 -719
  326. tace-0.2.0/tace/models/_e3nn/inter.py +0 -683
  327. tace-0.2.0/tace/models/_e3nn/layer_norm.py +0 -151
  328. tace-0.2.0/tace/models/_e3nn/node.py +0 -238
  329. tace-0.2.0/tace/models/_e3nn/paths.py +0 -75
  330. tace-0.2.0/tace/models/_e3nn/prod.py +0 -451
  331. tace-0.2.0/tace/models/_e3nn/readout.py +0 -200
  332. tace-0.2.0/tace/models/_e3nn/representation.py +0 -499
  333. tace-0.2.0/tace/models/_e3nn/symmetric_contraction.py +0 -566
  334. tace-0.2.0/tace/models/_eqt/__init__.py +0 -4
  335. tace-0.2.0/tace/models/_oeq/__init__.py +0 -6
  336. tace-0.2.0/tace/models/blocks.py +0 -467
  337. tace-0.2.0/tace/models/layout.py +0 -84
  338. tace-0.2.0/tace/models/legacy.py +0 -2171
  339. tace-0.2.0/tace/models/mlp.py +0 -283
  340. tace-0.2.0/tace/models/precision.py +0 -82
  341. tace-0.2.0/tace/models/s2.py +0 -1006
  342. tace-0.2.0/tace/models/so2/__init__.py +0 -25
  343. tace-0.2.0/tace/models/so2/blocks.py +0 -1122
  344. tace-0.2.0/tace/models/so2/rotation_matrix.py +0 -84
  345. tace-0.2.0/tace/models/so2/utils.py +0 -196
  346. tace-0.2.0/tace/models/so2/wigner.py +0 -365
  347. tace-0.2.0/tace/models/triton_ops/__init__.py +0 -9
  348. tace-0.2.0/tace/models/triton_ops/uu_so2_scatter.py +0 -837
  349. tace-0.2.0/tace/utils/loss/uncertainty.py +0 -89
  350. tace-0.2.0/tace/utils/optimizer/soap.py +0 -1
  351. tace-0.2.0/test/test_aoti_single_system.py +0 -207
  352. tace-0.2.0/test/test_neighbour_list.py +0 -188
  353. tace-0.2.0/test/test_solid_harmonics.py +0 -79
  354. tace-0.2.0/test/test_wigner6j.py +0 -214
  355. {tace-0.2.0 → tace-0.2.2}/LICENSE.md +0 -0
  356. {tace-0.2.0 → tace-0.2.2}/TACE.egg-info/dependency_links.txt +0 -0
  357. {tace-0.2.0 → tace-0.2.2}/setup.cfg +0 -0
  358. {tace-0.2.0 → tace-0.2.2}/tace/dataset/__init__.py +0 -0
  359. {tace-0.2.0 → tace-0.2.2}/tace/interface/__init__.py +0 -0
  360. {tace-0.2.0 → tace-0.2.2}/tace/interface/lammps/__init__.py +0 -0
  361. {tace-0.2.0 → tace-0.2.2}/tace/interface/lammps/mliap.py +0 -0
  362. {tace-0.2.0 → tace-0.2.2}/tace/interface/openmm/__init__.py +0 -0
  363. {tace-0.2.0 → tace-0.2.2}/tace/interface/uspex26/__init__.py +0 -0
  364. {tace-0.2.0 → tace-0.2.2}/tace/models/compile/__init__.py +2 -2
  365. {tace-0.2.0/tace/models/_cue → tace-0.2.2/tace/models/cue}/_tp_uuu.py +0 -0
  366. {tace-0.2.0/tace/models/_eqt → tace-0.2.2/tace/models/eqt}/equitorch/__init__.py +0 -0
  367. {tace-0.2.0/tace/models/_eqt → tace-0.2.2/tace/models/eqt}/equitorch/nn/__init__.py +0 -0
  368. {tace-0.2.0 → tace-0.2.2}/tace/scripts/__init__.py +0 -0
  369. {tace-0.2.0 → tace-0.2.2}/tace/utils/mask_metrics.py +0 -0
@@ -1,12 +1,12 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: TACE
3
- Version: 0.2.0
3
+ Version: 0.2.2
4
4
  Summary: TACE - ST and ICT based equivariant atomistic model.
5
5
  Author-email: Zemin Xu <xv_chana@163.com>
6
6
  Project-URL: Homepage, https://github.com/xvzemin/tace
7
7
  Requires-Python: >=3.9
8
8
  License-File: LICENSE.md
9
- Requires-Dist: torch<=2.13,>=2.4
9
+ Requires-Dist: torch>=2.4
10
10
  Requires-Dist: torch_geometric>=2.4
11
11
  Requires-Dist: lightning
12
12
  Requires-Dist: omegaconf
@@ -14,16 +14,27 @@ Requires-Dist: hydra-core>=1.3
14
14
  Requires-Dist: matscipy
15
15
  Requires-Dist: ase
16
16
  Requires-Dist: numpy
17
- Requires-Dist: e3nn
17
+ Requires-Dist: e3nn>=0.4.4
18
18
  Requires-Dist: scipy
19
19
  Requires-Dist: configargparse
20
20
  Requires-Dist: pandas
21
+ Requires-Dist: matplotlib
21
22
  Requires-Dist: lmdb
22
23
  Requires-Dist: huggingface_hub
24
+ Provides-Extra: dev
25
+ Requires-Dist: pytest; extra == "dev"
26
+ Requires-Dist: ruff<0.17,>=0.16; extra == "dev"
23
27
  Provides-Extra: oeq
24
28
  Requires-Dist: openequivariance>=0.5.4; extra == "oeq"
29
+ Provides-Extra: eqx
30
+ Requires-Dist: ninja; extra == "eqx"
31
+ Requires-Dist: filelock; extra == "eqx"
25
32
  Provides-Extra: torchsim
26
33
  Requires-Dist: torch-sim-atomistic>=0.6.1; extra == "torchsim"
34
+ Provides-Extra: nvalchemi
35
+ Requires-Dist: nvalchemi-toolkit>=0.1.0; extra == "nvalchemi"
36
+ Requires-Dist: nvalchemi-toolkit-ops; extra == "nvalchemi"
37
+ Provides-Extra: scf
27
38
  Provides-Extra: cueq12
28
39
  Requires-Dist: cuequivariance; extra == "cueq12"
29
40
  Requires-Dist: cuequivariance-torch; extra == "cueq12"
tace-0.2.2/README.md ADDED
@@ -0,0 +1,217 @@
1
+ [![TACE](https://img.shields.io/pypi/v/tace?style=for-the-badge&label=TACE)](https://pypi.org/project/tace/)
2
+ [![Docs](https://img.shields.io/readthedocs/tace?style=for-the-badge&label=docs)](https://tace.readthedocs.io/en/latest/)
3
+ [![License](https://img.shields.io/badge/License-MIT-yellow.svg?style=for-the-badge)](https://opensource.org/licenses/MIT)
4
+ <!-- [![Matbench Discovery](https://img.shields.io/badge/Matbench%20Discovery-SOTA-brightgreen?style=for-the-badge)](https://matbench-discovery.materialsproject.org/) -->
5
+
6
+ # Tensor Atomic/Edge Cluster Expansion (TACE/TECE)
7
+
8
+ TACE is designed with physical priors and strong inductive biases to enhance extrapolation capability.
9
+ It performs Atomic Cluster Expansion and Edge Cluster Expansion based on spherical tensors
10
+ or irreducible Cartesian tensors, with an optional attention architecture.
11
+
12
+ ## O(3) Cartesian Architecture
13
+
14
+ <img src="fig/cartesian_arch.png" width="100%" align="center">
15
+
16
+ ```math
17
+ (A\otimes_k^\delta B)_{\boldsymbol i\boldsymbol j}
18
+ =\frac{1}{\sqrt{3^k}}\sum_{\boldsymbol a}
19
+ A_{\boldsymbol i\boldsymbol a}B_{\boldsymbol a\boldsymbol j},
20
+ \qquad l_3=l_1+l_2-2k,
21
+ ````
22
+
23
+ ```math
24
+ (A\otimes_k^\epsilon B)_{\boldsymbol i w\boldsymbol j}
25
+ =\frac{1}{\sqrt{2}}\frac{1}{\sqrt{3^k}}\sum_{\boldsymbol a,u,v}
26
+ A_{\boldsymbol i u\boldsymbol a}\epsilon_{wuv}
27
+ B_{\boldsymbol a v\boldsymbol j},
28
+ \qquad l_3=l_1+l_2-2k-1.
29
+ ````
30
+
31
+ ## Time-reversal O(3) Spherical/O(2)/Magnetic Architecture
32
+
33
+ The SO(2) implementation in TACE will be gradually replaced by an O(2) formulation
34
+ to ensure complete parity support for O(3). Backward compatibility will be maintained for the SO(2) implementation,
35
+ while backward compatibility for the O(2) implementation is not currently guaranteed.
36
+
37
+ <img src="fig/o2.svg" width="100%" align="center">
38
+
39
+ ```math
40
+ \left.(l,p,t)_{O(3)\times\mathbb Z_2^{\mathcal T}}
41
+ \right\downarrow_{O(2)\times\mathbb Z_2^{\mathcal T}}
42
+ =
43
+ \underbrace{\bigl(0,p(-1)^l,t\bigr)}_{\text{1D},\,m=0}
44
+ \oplus
45
+ \underbrace{\bigoplus_{m=1}^{l}(m,0,t)}_{\text{2D},\,m>0}.
46
+ ````
47
+
48
+ ## Documentation
49
+
50
+ [TACE DOCS](https://tace.readthedocs.io/en/latest/index.html)
51
+
52
+ ## Foundation Model
53
+
54
+ [HUGGING FACE](https://huggingface.co/xvzemin/tace-foundations/tree/main)
55
+
56
+ [FOUNDATION MODEL DOCS](https://tace.readthedocs.io/en/latest/guide/foundation.html)
57
+
58
+ ## Install, Train and Tutorial
59
+
60
+ The docs contain a complete tutorial.
61
+
62
+ We also provide complete input files and example scripts for training, ASE,
63
+ TorchSim, LAMMPS, and other workflows in the
64
+ [TACE examples](https://github.com/xvzemin/tace/tree/main/example).
65
+
66
+ ```bash
67
+ # Minimal install and training example
68
+ git clone https://github.com/xvzemin/tace
69
+ cd tace
70
+ pip install . # or pip install tace
71
+ cd example/train
72
+ tace-train -cn tace.yaml
73
+ ```
74
+
75
+ ## Fine-tuning
76
+
77
+ - ✅ Full-parameter.
78
+
79
+ - ✅ Freeze-parameter.
80
+
81
+ - ✅ LoRA.
82
+
83
+ ## Overview
84
+
85
+ Currently, the officially supported properties include:
86
+
87
+ - Energy
88
+ - Forces (conservative | direct)
89
+ - Hessian (conservative, predict only)
90
+ - Stress (conservative | direct)
91
+ - Virials (conservative | direct)
92
+ - Charges (lagrangian or uniform_distribution)
93
+ - Dipole moment (conservative | direct)
94
+ - Polarization (conservative, multi-value for PBC systems)
95
+ - Polarizability (conservative | direct)
96
+ - Born effective charges (conservative, under electric field)
97
+ - Atomic stresses (conservative, predict only)
98
+ - Atomic virials (conservative, predict only)
99
+ - Absolute final collinear magmoms
100
+ - Collinear magnetic forces
101
+ - Noncollinear magnetic forces
102
+
103
+ ## Plugins
104
+
105
+ TACE currently supports the following plugin:
106
+
107
+ - **mTACE** (Magnetic, with and without Spin-Orbit Coupling)
108
+ - **TACE-LES** (Latent Ewald Summation)
109
+ - **TACE-QEq** (Lagrangian, under reconstruction)
110
+
111
+ ## Interfaces
112
+
113
+ - ✅ Supports integration with [ASE Calculator](https://wiki.fysik.dtu.dk/ase/).
114
+
115
+ - ✅ Supports integration with [LAMMPS-ML-IAP](https://github.com/lammps/lammps).
116
+
117
+ - ✅ Supports integration with [TorchSim](https://torchsim.github.io/torch-sim/).
118
+
119
+ - ✅ Supports integration with [NVIDIA NValCHEMI](https://github.com/NVIDIA/nvalchemi-toolkit).
120
+
121
+ - ✅ Supports integration with [OpenMM-ML](https://github.com/openmm/openmm-ml) (OpenMM-ML -> ASE -> TACE).
122
+
123
+ - ✅ Supports integration with [USPEX](https://uspex-team.org/)
124
+ (USPEX -> LAMMPS-ML-IAP -> TACE) (Python=3.9).
125
+
126
+ ## Contact
127
+
128
+ For bugs or feature requests, please use the
129
+ [TACE issue](https://github.com/xvzemin/tace/issues).
130
+
131
+ <!-- For usage discussions, you can join the TACE community through QQ or Discord.
132
+
133
+ <table>
134
+ <tr>
135
+ <th>QQ</th>
136
+ <th>Discord</th>
137
+ </tr>
138
+ <tr>
139
+ <td align="center">
140
+ <img src="fig/qq.jpg" alt="TACE QQ QR code" width="260">
141
+ </td>
142
+ <td align="center">
143
+ <img src="fig/discord.jpg" alt="TACE Discord QR code" width="260">
144
+ </td>
145
+ </tr>
146
+ </table> -->
147
+
148
+ ## Citing
149
+
150
+ If you use TACE, please cite our papers:
151
+
152
+ ```bibtex
153
+ @misc{xu2026spectralspatialtensoratomiccluster,
154
+ title={Spectral/Spatial Tensor Atomic Cluster Expansion with Universal Embeddings in Cartesian Space},
155
+ author={Zemin Xu and Wenbo Xie and P. Hu},
156
+ year={2026},
157
+ eprint={2509.14961},
158
+ archivePrefix={arXiv},
159
+ primaryClass={stat.ML},
160
+ url={https://arxiv.org/abs/2509.14961},
161
+ }
162
+
163
+ @misc{xu2026edgeclusterexpansionradial,
164
+ title={Edge Cluster Expansion with Radial Rotary Attention for Interatomic Potentials},
165
+ author={Zemin Xu and Wenbo Xie and P. Hu},
166
+ year={2026},
167
+ eprint={2607.10664},
168
+ archivePrefix={arXiv},
169
+ primaryClass={stat.ML},
170
+ url={https://arxiv.org/abs/2607.10664},
171
+ }
172
+ ```
173
+
174
+ If you use Local O(2) Frame, please cite our papers:
175
+
176
+ ```bibtex
177
+ @misc{xu2026lookingglassefficientparitycompletelearning,
178
+ title={Through the Looking-Glass: Efficient Parity-Complete Learning via Local $O(2)$ Frames},
179
+ author={Zemin Xu and Wenbo Xie},
180
+ year={2026},
181
+ eprint={2608.16592},
182
+ archivePrefix={arXiv},
183
+ primaryClass={physics.chem-ph},
184
+ url={https://arxiv.org/abs/2608.16592},
185
+ }
186
+ ```
187
+
188
+ If you use cartnn, Cartesian-3j, cMACE, cNequIP, cAllegro, please cite our papers:
189
+
190
+ ```bibtex
191
+ @inproceedings{xu2026a,
192
+ title={A Cartesian-3j Framework for Machine Learning Interatomic Potentials},
193
+ author={Zemin Xu and Chenyu Wu and Wenbo Xie and Peijun Hu},
194
+ booktitle={Forty-third International Conference on Machine Learning},
195
+ year={2026},
196
+ url={https://openreview.net/forum?id=9ZWK6gneWq}
197
+ }
198
+ ```
199
+
200
+ ## Development
201
+
202
+ Install the development tools and run the same formatting and lint checks used by
203
+ pre-commit:
204
+
205
+ ```bash
206
+ pip install -e ".[dev]"
207
+ ruff check --fix tace
208
+ ruff format tace
209
+ ruff check tace
210
+ ruff format --check tace
211
+ ```
212
+
213
+ ## License
214
+
215
+ The TACE code is published and distributed under the MIT License.
216
+ EquivariantX (`eqx`) is licensed separately under
217
+ [Creative Commons Attribution 4.0 International (CC BY 4.0)](eqx/LICENSE.md).
@@ -1,12 +1,12 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: TACE
3
- Version: 0.2.0
3
+ Version: 0.2.2
4
4
  Summary: TACE - ST and ICT based equivariant atomistic model.
5
5
  Author-email: Zemin Xu <xv_chana@163.com>
6
6
  Project-URL: Homepage, https://github.com/xvzemin/tace
7
7
  Requires-Python: >=3.9
8
8
  License-File: LICENSE.md
9
- Requires-Dist: torch<=2.13,>=2.4
9
+ Requires-Dist: torch>=2.4
10
10
  Requires-Dist: torch_geometric>=2.4
11
11
  Requires-Dist: lightning
12
12
  Requires-Dist: omegaconf
@@ -14,16 +14,27 @@ Requires-Dist: hydra-core>=1.3
14
14
  Requires-Dist: matscipy
15
15
  Requires-Dist: ase
16
16
  Requires-Dist: numpy
17
- Requires-Dist: e3nn
17
+ Requires-Dist: e3nn>=0.4.4
18
18
  Requires-Dist: scipy
19
19
  Requires-Dist: configargparse
20
20
  Requires-Dist: pandas
21
+ Requires-Dist: matplotlib
21
22
  Requires-Dist: lmdb
22
23
  Requires-Dist: huggingface_hub
24
+ Provides-Extra: dev
25
+ Requires-Dist: pytest; extra == "dev"
26
+ Requires-Dist: ruff<0.17,>=0.16; extra == "dev"
23
27
  Provides-Extra: oeq
24
28
  Requires-Dist: openequivariance>=0.5.4; extra == "oeq"
29
+ Provides-Extra: eqx
30
+ Requires-Dist: ninja; extra == "eqx"
31
+ Requires-Dist: filelock; extra == "eqx"
25
32
  Provides-Extra: torchsim
26
33
  Requires-Dist: torch-sim-atomistic>=0.6.1; extra == "torchsim"
34
+ Provides-Extra: nvalchemi
35
+ Requires-Dist: nvalchemi-toolkit>=0.1.0; extra == "nvalchemi"
36
+ Requires-Dist: nvalchemi-toolkit-ops; extra == "nvalchemi"
37
+ Provides-Extra: scf
27
38
  Provides-Extra: cueq12
28
39
  Requires-Dist: cuequivariance; extra == "cueq12"
29
40
  Requires-Dist: cuequivariance-torch; extra == "cueq12"
@@ -0,0 +1,308 @@
1
+ LICENSE.md
2
+ README.md
3
+ pyproject.toml
4
+ TACE.egg-info/PKG-INFO
5
+ TACE.egg-info/SOURCES.txt
6
+ TACE.egg-info/dependency_links.txt
7
+ TACE.egg-info/entry_points.txt
8
+ TACE.egg-info/requires.txt
9
+ TACE.egg-info/top_level.txt
10
+ eqx/LICENSE.md
11
+ eqx/README.md
12
+ eqx/__init__.py
13
+ eqx/_layout.py
14
+ eqx/utils.py
15
+ eqx/ace/__init__.py
16
+ eqx/ace/contraction.py
17
+ eqx/ace/cuda.py
18
+ eqx/ace/tace.py
19
+ eqx/co2/__init__.py
20
+ eqx/co2/basis.py
21
+ eqx/co2/cartesian_harmonics.py
22
+ eqx/co2/gate.py
23
+ eqx/co2/irreps.py
24
+ eqx/co2/linear.py
25
+ eqx/co2/o3_tensor_product.py
26
+ eqx/co2/restriction.py
27
+ eqx/co2/spherical.py
28
+ eqx/co2/tensor_product.py
29
+ eqx/co3/__init__.py
30
+ eqx/co3/basis.py
31
+ eqx/co3/cartesian_harmonics.py
32
+ eqx/co3/gate.py
33
+ eqx/co3/irreps.py
34
+ eqx/co3/linear.py
35
+ eqx/co3/symmetric.py
36
+ eqx/co3/tensor_product.py
37
+ eqx/conv/__init__.py
38
+ eqx/conv/angular.py
39
+ eqx/conv/attention.py
40
+ eqx/conv/codegen.py
41
+ eqx/conv/contraction.py
42
+ eqx/conv/edge.py
43
+ eqx/conv/graph.py
44
+ eqx/conv/program.py
45
+ eqx/conv/radial.py
46
+ eqx/conv/co3/__init__.py
47
+ eqx/conv/co3/convolution.py
48
+ eqx/conv/co3/linear.py
49
+ eqx/conv/co3/polynomials.py
50
+ eqx/conv/o2_o3/__init__.py
51
+ eqx/conv/o2_o3/autotune.py
52
+ eqx/conv/o2_o3/codegen.py
53
+ eqx/conv/o2_o3/convolution.py
54
+ eqx/conv/o2_o3/cuda.py
55
+ eqx/conv/o2_o3/direction_codegen.py
56
+ eqx/conv/o2_o3/geometry.py
57
+ eqx/conv/o2_o3/schedule.py
58
+ eqx/conv/o2_o3/transverse.py
59
+ eqx/conv/o3/__init__.py
60
+ eqx/conv/o3/codegen.py
61
+ eqx/conv/o3/convolution.py
62
+ eqx/conv/o3/cuda.py
63
+ eqx/conv/o3/harmonics.py
64
+ eqx/conv/uu_o2/__init__.py
65
+ eqx/conv/uu_o2/convolution.py
66
+ eqx/conv/uv_o2/__init__.py
67
+ eqx/conv/uv_o2/convolution.py
68
+ eqx/conv/uv_o2/transverse.py
69
+ eqx/kernels/__init__.py
70
+ eqx/kernels/channel_product.py
71
+ eqx/kernels/codegen.py
72
+ eqx/kernels/cuda.py
73
+ eqx/kernels/cuda_graph.py
74
+ eqx/kernels/layout.py
75
+ eqx/kernels/quaternion.py
76
+ eqx/kernels/recompute.py
77
+ eqx/kernels/rotary.py
78
+ eqx/kernels/rotation.py
79
+ eqx/kernels/wigner.py
80
+ eqx/kernels/csrc/runtime.cpp
81
+ eqx/models/__init__.py
82
+ eqx/models/convolution.py
83
+ eqx/models/equflash/__init__.py
84
+ eqx/models/equflash/conversion.py
85
+ eqx/models/mace/__init__.py
86
+ eqx/models/mace/conversion.py
87
+ eqx/models/nequip/__init__.py
88
+ eqx/models/nequip/conversion.py
89
+ eqx/models/prophet/__init__.py
90
+ eqx/models/prophet/conversion.py
91
+ eqx/models/sevennet/__init__.py
92
+ eqx/models/sevennet/conversion.py
93
+ eqx/models/tace/__init__.py
94
+ eqx/models/tace/tece_oam_rra/__init__.py
95
+ eqx/models/tace/tece_oam_rra/bilinear_contraction.py
96
+ eqx/models/tace/tece_oam_rra/bilinear_cuda.py
97
+ eqx/models/tace/tece_oam_rra/cuda.py
98
+ eqx/models/tace/tece_oam_rra/execution.py
99
+ eqx/models/tace/tece_oam_rra/interaction.py
100
+ eqx/models/tace/tece_oam_rra/product.py
101
+ eqx/models/tace/tece_oam_rra/program.py
102
+ eqx/o2/__init__.py
103
+ eqx/o2/_clebsch_gordan.py
104
+ eqx/o2/_layout.py
105
+ eqx/o2/asymmetric_contraction.py
106
+ eqx/o2/circular_harmonics.py
107
+ eqx/o2/gate.py
108
+ eqx/o2/irreps.py
109
+ eqx/o2/linear.py
110
+ eqx/o2/local_frame.py
111
+ eqx/o2/o3_tensor_product.py
112
+ eqx/o2/rotation_matrix.py
113
+ eqx/o2/tensor_product.py
114
+ eqx/o2/wigner.py
115
+ eqx/o3/__init__.py
116
+ eqx/o3/contraction.py
117
+ eqx/o3/cuda.py
118
+ eqx/o3/gate.py
119
+ eqx/o3/linear.py
120
+ tace/__init__.py
121
+ tace/dataset/__init__.py
122
+ tace/dataset/augmentation.py
123
+ tace/dataset/dataloader.py
124
+ tace/dataset/datamodule.py
125
+ tace/dataset/element.py
126
+ tace/dataset/graph.py
127
+ tace/dataset/neighbour_list.py
128
+ tace/dataset/quantity.py
129
+ tace/dataset/read.py
130
+ tace/dataset/sampler.py
131
+ tace/dataset/split.py
132
+ tace/dataset/statistics.py
133
+ tace/dataset/utils.py
134
+ tace/foundations/__init__.py
135
+ tace/foundations/download_link.py
136
+ tace/foundations/u_shift.py
137
+ tace/interface/__init__.py
138
+ tace/interface/ase/__init__.py
139
+ tace/interface/ase/calculator.py
140
+ tace/interface/ase/general_calculator.py
141
+ tace/interface/ase/optimizer.py
142
+ tace/interface/lammps/__init__.py
143
+ tace/interface/lammps/mliap.py
144
+ tace/interface/nvalchemi/__init__.py
145
+ tace/interface/nvalchemi/wrapper.py
146
+ tace/interface/openmm/__init__.py
147
+ tace/interface/torchsim/__init__.py
148
+ tace/interface/torchsim/torchsim.py
149
+ tace/interface/uspex26/__init__.py
150
+ tace/interface/uspex26/relax.py
151
+ tace/lightning/__init__.py
152
+ tace/lightning/lit_model.py
153
+ tace/lightning/lora.py
154
+ tace/lightning/loss_skip.py
155
+ tace/lightning/torch_model.py
156
+ tace/lightning/trainer.py
157
+ tace/lightning/u_shift.py
158
+ tace/models/__init__.py
159
+ tace/models/adapter.py
160
+ tace/models/angular.py
161
+ tace/models/blocks.py
162
+ tace/models/ictd.py
163
+ tace/models/kspace.py
164
+ tace/models/lammps.py
165
+ tace/models/layout.py
166
+ tace/models/linear.py
167
+ tace/models/mlp.py
168
+ tace/models/normalizer.py
169
+ tace/models/radial.py
170
+ tace/models/s2.py
171
+ tace/models/softmax.py
172
+ tace/models/time_reversal.py
173
+ tace/models/utils.py
174
+ tace/models/zemin.py
175
+ tace/models/_e3nn/__init__.py
176
+ tace/models/_e3nn/base.py
177
+ tace/models/_e3nn/basis_change.py
178
+ tace/models/_e3nn/default.py
179
+ tace/models/_e3nn/dropout.py
180
+ tace/models/_e3nn/edge.py
181
+ tace/models/_e3nn/fused.py
182
+ tace/models/_e3nn/inter.py
183
+ tace/models/_e3nn/layer_norm.py
184
+ tace/models/_e3nn/les.py
185
+ tace/models/_e3nn/magnetic.py
186
+ tace/models/_e3nn/node.py
187
+ tace/models/_e3nn/nonlinear.py
188
+ tace/models/_e3nn/o2.py
189
+ tace/models/_e3nn/paths.py
190
+ tace/models/_e3nn/prod.py
191
+ tace/models/_e3nn/readout.py
192
+ tace/models/_e3nn/representation.py
193
+ tace/models/_e3nn/residual.py
194
+ tace/models/_e3nn/scatter_norm.py
195
+ tace/models/_e3nn/symmetric_contraction.py
196
+ tace/models/_e3nn/tace.py
197
+ tace/models/_e3nn/tece_oam_rra.py
198
+ tace/models/_e3nn/ue.py
199
+ tace/models/_e3nn/wigner6j.py
200
+ tace/models/compile/__init__.py
201
+ tace/models/compile/aot.py
202
+ tace/models/compile/compile.py
203
+ tace/models/compile/tace.py
204
+ tace/models/compile/wrapper.py
205
+ tace/models/cue/__init__.py
206
+ tace/models/cue/_tp_scatter.py
207
+ tace/models/cue/_tp_uuu.py
208
+ tace/models/cue/paths.py
209
+ tace/models/eqt/__init__.py
210
+ tace/models/eqt/_tp_uuu.py
211
+ tace/models/eqt/equitorch/__init__.py
212
+ tace/models/eqt/equitorch/irreps/__init__.py
213
+ tace/models/eqt/equitorch/nn/__init__.py
214
+ tace/models/eqt/equitorch/nn/sparse_product.py
215
+ tace/models/eqt/equitorch/nn/tensor_products.py
216
+ tace/models/eqt/equitorch/structs/__init__.py
217
+ tace/models/oeq/__init__.py
218
+ tace/models/oeq/_tp_scatter.py
219
+ tace/models/scf/README.md
220
+ tace/models/scf/__init__.py
221
+ tace/models/scf/electrostatics.py
222
+ tace/models/scf/energy_functional.py
223
+ tace/models/scf/fixed_point.py
224
+ tace/models/scf/local.py
225
+ tace/models/scf/loss.py
226
+ tace/models/scf/model.py
227
+ tace/models/scf/qeq.py
228
+ tace/models/scf/readout.py
229
+ tace/models/scf/state.py
230
+ tace/models/scf/utils.py
231
+ tace/models/scf/longrange/LICENSE.md
232
+ tace/models/scf/longrange/NOTICE.md
233
+ tace/models/scf/longrange/__init__.py
234
+ tace/models/scf/longrange/__version__.py
235
+ tace/models/scf/longrange/energy.py
236
+ tace/models/scf/longrange/features.py
237
+ tace/models/scf/longrange/gto_utils.py
238
+ tace/models/scf/longrange/kspace.py
239
+ tace/models/scf/longrange/realspace_electrostatics.py
240
+ tace/models/scf/longrange/realspace_grid_integrals.py
241
+ tace/models/scf/longrange/slabs.py
242
+ tace/models/scf/longrange/utils.py
243
+ tace/scripts/__init__.py
244
+ tace/scripts/average.py
245
+ tace/scripts/clean.py
246
+ tace/scripts/convert_cgtp.py
247
+ tace/scripts/convert_lora.py
248
+ tace/scripts/eval.py
249
+ tace/scripts/export_eval.py
250
+ tace/scripts/export_lammps.py
251
+ tace/scripts/export_train.py
252
+ tace/scripts/finetune.py
253
+ tace/scripts/plot_diatom.py
254
+ tace/scripts/split.py
255
+ tace/scripts/train.py
256
+ tace/scripts/update.py
257
+ tace/scripts/utils/__init__.py
258
+ tace/scripts/utils/check_equi.py
259
+ tace/scripts/utils/check_soc.py
260
+ tace/scripts/utils/convert_to_xyz.py
261
+ tace/utils/__init__.py
262
+ tace/utils/_global.py
263
+ tace/utils/callbacks.py
264
+ tace/utils/cfg.py
265
+ tace/utils/ema.py
266
+ tace/utils/env.py
267
+ tace/utils/hydra_resolver.py
268
+ tace/utils/logger.py
269
+ tace/utils/mask_metrics.py
270
+ tace/utils/metrics.py
271
+ tace/utils/metrics_bake.py
272
+ tace/utils/spectra.py
273
+ tace/utils/strategy.py
274
+ tace/utils/torch_scatter.py
275
+ tace/utils/units.py
276
+ tace/utils/utils.py
277
+ tace/utils/loss/__init__.py
278
+ tace/utils/loss/common.py
279
+ tace/utils/loss/dens.py
280
+ tace/utils/loss/huber_fn.py
281
+ tace/utils/loss/l2mae_fn.py
282
+ tace/utils/loss/mae_fn.py
283
+ tace/utils/loss/mse_fn.py
284
+ tace/utils/loss/normal.py
285
+ tace/utils/loss/registry.py
286
+ tace/utils/loss/special_fn.py
287
+ tace/utils/loss/uncertainty.py
288
+ tace/utils/lr_scheduler/__init__.py
289
+ tace/utils/lr_scheduler/warmup.py
290
+ tace/utils/lr_scheduler/wsd.py
291
+ tace/utils/optimizer/__init__.py
292
+ tace/utils/optimizer/hybrid_muon.py
293
+ tace/utils/optimizer/soap.py
294
+ tests/test_cartesian.py
295
+ tests/test_compile.py
296
+ tests/test_eqt.py
297
+ tests/test_eqx_integration.py
298
+ tests/test_les.py
299
+ tests/test_local_o2.py
300
+ tests/test_loss.py
301
+ tests/test_magnetic_optimizer.py
302
+ tests/test_product.py
303
+ tests/test_radial.py
304
+ tests/test_scf.py
305
+ tests/test_statistics.py
306
+ tests/test_time_reversal.py
307
+ tests/test_torchsim.py
308
+ tests/test_wigner6j.py
@@ -2,12 +2,14 @@
2
2
  tace-average = tace.scripts.average:main
3
3
  tace-clean = tace.scripts.clean:main
4
4
  tace-compile = tace.scripts.export_eval:main
5
- tace-convert = tace.scripts.convert:main
5
+ tace-convert-cgtp = tace.scripts.convert_cgtp:main
6
+ tace-convert-lora = tace.scripts.convert_lora:main
6
7
  tace-eval = tace.scripts.eval:main
7
8
  tace-export-eval = tace.scripts.export_eval:main
8
9
  tace-export-lammps = tace.scripts.export_lammps:main
9
10
  tace-export-train = tace.scripts.export_train:main
10
11
  tace-finetune = tace.scripts.finetune:main
12
+ tace-plot-diatom = tace.scripts.plot_diatom:main
11
13
  tace-split = tace.scripts.split:main
12
14
  tace-train = tace.scripts.train:main
13
15
  tace-update = tace.scripts.update:main
@@ -1,4 +1,4 @@
1
- torch<=2.13,>=2.4
1
+ torch>=2.4
2
2
  torch_geometric>=2.4
3
3
  lightning
4
4
  omegaconf
@@ -6,10 +6,11 @@ hydra-core>=1.3
6
6
  matscipy
7
7
  ase
8
8
  numpy
9
- e3nn
9
+ e3nn>=0.4.4
10
10
  scipy
11
11
  configargparse
12
12
  pandas
13
+ matplotlib
13
14
  lmdb
14
15
  huggingface_hub
15
16
 
@@ -26,6 +27,14 @@ cuequivariance-ops-torch-cu13
26
27
  [d3]
27
28
  torch_dftd
28
29
 
30
+ [dev]
31
+ pytest
32
+ ruff<0.17,>=0.16
33
+
34
+ [eqx]
35
+ ninja
36
+ filelock
37
+
29
38
  [lammps12]
30
39
  cython
31
40
  cupy-cuda12x
@@ -39,6 +48,10 @@ ase
39
48
  matscipy
40
49
  vesin
41
50
 
51
+ [nvalchemi]
52
+ nvalchemi-toolkit>=0.1.0
53
+ nvalchemi-toolkit-ops
54
+
42
55
  [oeq]
43
56
  openequivariance>=0.5.4
44
57
 
@@ -47,5 +60,7 @@ openmm
47
60
  openmm-ml
48
61
  openmmtorch
49
62
 
63
+ [scf]
64
+
50
65
  [torchsim]
51
66
  torch-sim-atomistic>=0.6.1
@@ -1 +1,2 @@
1
+ eqx
1
2
  tace