k3-node 1.0.0__py3-none-any.whl

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 (459) hide show
  1. k3_node/__init__.py +122 -0
  2. k3_node/applications/__init__.py +17 -0
  3. k3_node/applications/bio/__init__.py +21 -0
  4. k3_node/applications/chemistry/__init__.py +155 -0
  5. k3_node/applications/materials/__init__.py +127 -0
  6. k3_node/applications/materials/basis.py +449 -0
  7. k3_node/applications/materials/chgnet.py +360 -0
  8. k3_node/applications/materials/core.py +351 -0
  9. k3_node/applications/materials/grace.py +246 -0
  10. k3_node/applications/materials/io.py +230 -0
  11. k3_node/applications/materials/m3gnet.py +462 -0
  12. k3_node/applications/materials/megnet.py +395 -0
  13. k3_node/applications/materials/qet.py +220 -0
  14. k3_node/applications/materials/readout.py +235 -0
  15. k3_node/applications/materials/so3net.py +234 -0
  16. k3_node/applications/materials/tensornet.py +381 -0
  17. k3_node/applications/materials/test_materials.py +167 -0
  18. k3_node/applications/materials/wrappers.py +95 -0
  19. k3_node/data/__init__.py +47 -0
  20. k3_node/data/batch.py +102 -0
  21. k3_node/data/collate.py +282 -0
  22. k3_node/data/data.py +532 -0
  23. k3_node/data/database.py +154 -0
  24. k3_node/data/dataset.py +182 -0
  25. k3_node/data/download.py +49 -0
  26. k3_node/data/extract.py +45 -0
  27. k3_node/data/feature_store.py +70 -0
  28. k3_node/data/graph_store.py +92 -0
  29. k3_node/data/hetero_data.py +374 -0
  30. k3_node/data/hypergraph_data.py +59 -0
  31. k3_node/data/in_memory_dataset.py +177 -0
  32. k3_node/data/makedirs.py +7 -0
  33. k3_node/data/on_disk_dataset.py +77 -0
  34. k3_node/data/separate.py +115 -0
  35. k3_node/data/storage.py +593 -0
  36. k3_node/data/temporal.py +154 -0
  37. k3_node/data/test_batch.py +67 -0
  38. k3_node/data/test_data.py +68 -0
  39. k3_node/data/test_dataset_and_stores.py +111 -0
  40. k3_node/data/test_hetero_data.py +33 -0
  41. k3_node/data/test_temporal_and_hyper.py +32 -0
  42. k3_node/data/view.py +43 -0
  43. k3_node/datasets/__init__.py +88 -0
  44. k3_node/datasets/actor.py +101 -0
  45. k3_node/datasets/airports.py +84 -0
  46. k3_node/datasets/amazon.py +66 -0
  47. k3_node/datasets/ba2motif_dataset.py +73 -0
  48. k3_node/datasets/ba_shapes.py +81 -0
  49. k3_node/datasets/bitcoin_otc.py +77 -0
  50. k3_node/datasets/citation_full.py +81 -0
  51. k3_node/datasets/coauthor.py +66 -0
  52. k3_node/datasets/dblp.py +106 -0
  53. k3_node/datasets/digits.py +63 -0
  54. k3_node/datasets/email_eu_core.py +60 -0
  55. k3_node/datasets/entities.py +158 -0
  56. k3_node/datasets/explainer_dataset.py +101 -0
  57. k3_node/datasets/facebook.py +51 -0
  58. k3_node/datasets/fake.py +256 -0
  59. k3_node/datasets/freebase.py +90 -0
  60. k3_node/datasets/geometric_shapes.py +69 -0
  61. k3_node/datasets/github.py +51 -0
  62. k3_node/datasets/graph_generator/__init__.py +6 -0
  63. k3_node/datasets/graph_generator/ba_graph.py +20 -0
  64. k3_node/datasets/graph_generator/base.py +29 -0
  65. k3_node/datasets/graph_generator/er_graph.py +21 -0
  66. k3_node/datasets/icews.py +58 -0
  67. k3_node/datasets/imdb.py +96 -0
  68. k3_node/datasets/jodie.py +56 -0
  69. k3_node/datasets/karate.py +56 -0
  70. k3_node/datasets/lastfm_asia.py +51 -0
  71. k3_node/datasets/mesh_correspondence.py +50 -0
  72. k3_node/datasets/molecule_net.py +148 -0
  73. k3_node/datasets/motif_generator/__init__.py +7 -0
  74. k3_node/datasets/motif_generator/base.py +29 -0
  75. k3_node/datasets/motif_generator/custom.py +17 -0
  76. k3_node/datasets/motif_generator/cycle.py +25 -0
  77. k3_node/datasets/motif_generator/house.py +27 -0
  78. k3_node/datasets/movielens.py +55 -0
  79. k3_node/datasets/planetoid.py +137 -0
  80. k3_node/datasets/polblogs.py +63 -0
  81. k3_node/datasets/ppi.py +189 -0
  82. k3_node/datasets/qm7.py +65 -0
  83. k3_node/datasets/qm9.py +132 -0
  84. k3_node/datasets/reddit.py +121 -0
  85. k3_node/datasets/sbm_dataset.py +165 -0
  86. k3_node/datasets/seal.py +74 -0
  87. k3_node/datasets/shape_scenes.py +92 -0
  88. k3_node/datasets/test_datasets.py +322 -0
  89. k3_node/datasets/tu_dataset.py +131 -0
  90. k3_node/datasets/twitch.py +66 -0
  91. k3_node/datasets/webkb.py +102 -0
  92. k3_node/datasets/wikics.py +85 -0
  93. k3_node/datasets/word_net.py +184 -0
  94. k3_node/etl/__init__.py +37 -0
  95. k3_node/etl/encoders.py +248 -0
  96. k3_node/etl/graph_builders.py +270 -0
  97. k3_node/etl/relational_to_graph.py +201 -0
  98. k3_node/etl/table_to_graph.py +244 -0
  99. k3_node/etl/test_etl.py +318 -0
  100. k3_node/export/__init__.py +15 -0
  101. k3_node/export/cross_backend.py +172 -0
  102. k3_node/export/onnx_exporter.py +190 -0
  103. k3_node/export/runtime.py +254 -0
  104. k3_node/export/tensorrt_exporter.py +201 -0
  105. k3_node/export/test_export.py +337 -0
  106. k3_node/export/tflite_exporter.py +112 -0
  107. k3_node/hub/__init__.py +29 -0
  108. k3_node/hub/dataset_hub.py +242 -0
  109. k3_node/hub/hub_mixin.py +599 -0
  110. k3_node/hub/model_card.py +133 -0
  111. k3_node/hub/test_hub.py +419 -0
  112. k3_node/io/__init__.py +22 -0
  113. k3_node/io/fs.py +117 -0
  114. k3_node/io/npz.py +45 -0
  115. k3_node/io/off.py +29 -0
  116. k3_node/io/planetoid.py +98 -0
  117. k3_node/io/tu.py +137 -0
  118. k3_node/io/txt_array.py +58 -0
  119. k3_node/layers/__init__.py +14 -0
  120. k3_node/layers/aggr/__init__.py +70 -0
  121. k3_node/layers/aggr/attention.py +77 -0
  122. k3_node/layers/aggr/base.py +403 -0
  123. k3_node/layers/aggr/basic.py +412 -0
  124. k3_node/layers/aggr/deep_sets.py +65 -0
  125. k3_node/layers/aggr/deepsets.py +29 -0
  126. k3_node/layers/aggr/equilibrium.py +107 -0
  127. k3_node/layers/aggr/fused.py +43 -0
  128. k3_node/layers/aggr/gmt.py +89 -0
  129. k3_node/layers/aggr/gru.py +58 -0
  130. k3_node/layers/aggr/lcm.py +143 -0
  131. k3_node/layers/aggr/lstm.py +58 -0
  132. k3_node/layers/aggr/mlp.py +75 -0
  133. k3_node/layers/aggr/multi.py +154 -0
  134. k3_node/layers/aggr/patch_transformer.py +137 -0
  135. k3_node/layers/aggr/quantile.py +125 -0
  136. k3_node/layers/aggr/resolver.py +68 -0
  137. k3_node/layers/aggr/scaler.py +133 -0
  138. k3_node/layers/aggr/set2set.py +87 -0
  139. k3_node/layers/aggr/set_transformer.py +107 -0
  140. k3_node/layers/aggr/sort.py +68 -0
  141. k3_node/layers/aggr/test_aggr.py +337 -0
  142. k3_node/layers/aggr/utils.py +210 -0
  143. k3_node/layers/aggr/variance_preserving.py +54 -0
  144. k3_node/layers/attention/__init__.py +5 -0
  145. k3_node/layers/attention/pair_attention.py +448 -0
  146. k3_node/layers/attention/performer.py +187 -0
  147. k3_node/layers/attention/polynormer.py +160 -0
  148. k3_node/layers/attention/qformer.py +143 -0
  149. k3_node/layers/attention/sgformer.py +106 -0
  150. k3_node/layers/attention/test_attention.py +68 -0
  151. k3_node/layers/attention/test_pair_attention.py +91 -0
  152. k3_node/layers/conv/__init__.py +149 -0
  153. k3_node/layers/conv/agnn_conv.py +120 -0
  154. k3_node/layers/conv/antisymmetric_conv.py +94 -0
  155. k3_node/layers/conv/appnp.py +105 -0
  156. k3_node/layers/conv/appnp_conv.py +157 -0
  157. k3_node/layers/conv/arma_conv.py +231 -0
  158. k3_node/layers/conv/cg_conv.py +92 -0
  159. k3_node/layers/conv/cheb_conv.py +137 -0
  160. k3_node/layers/conv/cluster_gcn_conv.py +102 -0
  161. k3_node/layers/conv/conv.py +100 -0
  162. k3_node/layers/conv/crystal_conv.py +140 -0
  163. k3_node/layers/conv/cugraph.py +84 -0
  164. k3_node/layers/conv/diffusion_conv.py +144 -0
  165. k3_node/layers/conv/dir_gnn_conv.py +93 -0
  166. k3_node/layers/conv/dna_conv.py +192 -0
  167. k3_node/layers/conv/edge_conv.py +107 -0
  168. k3_node/layers/conv/eg_conv.py +155 -0
  169. k3_node/layers/conv/fa_conv.py +107 -0
  170. k3_node/layers/conv/feast_conv.py +126 -0
  171. k3_node/layers/conv/film_conv.py +143 -0
  172. k3_node/layers/conv/gat_conv.py +244 -0
  173. k3_node/layers/conv/gated_graph_conv.py +136 -0
  174. k3_node/layers/conv/gatv2_conv.py +205 -0
  175. k3_node/layers/conv/gcn.py +144 -0
  176. k3_node/layers/conv/gcn2_conv.py +126 -0
  177. k3_node/layers/conv/gcn_conv.py +135 -0
  178. k3_node/layers/conv/gen_conv.py +163 -0
  179. k3_node/layers/conv/general_conv.py +218 -0
  180. k3_node/layers/conv/gin_conv.py +218 -0
  181. k3_node/layers/conv/gmm_conv.py +172 -0
  182. k3_node/layers/conv/gps_conv.py +153 -0
  183. k3_node/layers/conv/graph_attention.py +262 -0
  184. k3_node/layers/conv/graph_conv.py +84 -0
  185. k3_node/layers/conv/gravnet_conv.py +93 -0
  186. k3_node/layers/conv/han_conv.py +175 -0
  187. k3_node/layers/conv/heat_conv.py +131 -0
  188. k3_node/layers/conv/hetero_conv.py +128 -0
  189. k3_node/layers/conv/hgt_conv.py +218 -0
  190. k3_node/layers/conv/hypergraph_conv.py +182 -0
  191. k3_node/layers/conv/le_conv.py +81 -0
  192. k3_node/layers/conv/lg_conv.py +58 -0
  193. k3_node/layers/conv/meshcnn_conv.py +84 -0
  194. k3_node/layers/conv/message_passing.py +451 -0
  195. k3_node/layers/conv/mf_conv.py +95 -0
  196. k3_node/layers/conv/mixhop_conv.py +108 -0
  197. k3_node/layers/conv/nn_conv.py +110 -0
  198. k3_node/layers/conv/pan_conv.py +100 -0
  199. k3_node/layers/conv/pdn_conv.py +109 -0
  200. k3_node/layers/conv/pna_conv.py +177 -0
  201. k3_node/layers/conv/point_conv.py +101 -0
  202. k3_node/layers/conv/point_gnn_conv.py +90 -0
  203. k3_node/layers/conv/point_transformer_conv.py +132 -0
  204. k3_node/layers/conv/ppf_conv.py +135 -0
  205. k3_node/layers/conv/ppnp.py +89 -0
  206. k3_node/layers/conv/res_gated_graph_conv.py +126 -0
  207. k3_node/layers/conv/rgat_conv.py +251 -0
  208. k3_node/layers/conv/rgcn_conv.py +321 -0
  209. k3_node/layers/conv/sage_conv.py +154 -0
  210. k3_node/layers/conv/sg_conv.py +96 -0
  211. k3_node/layers/conv/signed_conv.py +100 -0
  212. k3_node/layers/conv/simple_conv.py +75 -0
  213. k3_node/layers/conv/spline_conv.py +182 -0
  214. k3_node/layers/conv/ssg_conv.py +101 -0
  215. k3_node/layers/conv/supergat_conv.py +195 -0
  216. k3_node/layers/conv/tag_conv.py +98 -0
  217. k3_node/layers/conv/test_backend_consistency.py +164 -0
  218. k3_node/layers/conv/test_conv.py +176 -0
  219. k3_node/layers/conv/test_conv_pyg.py +566 -0
  220. k3_node/layers/conv/transformer_conv.py +168 -0
  221. k3_node/layers/conv/utils.py +403 -0
  222. k3_node/layers/conv/wl_conv.py +151 -0
  223. k3_node/layers/conv/x_conv.py +187 -0
  224. k3_node/layers/dense/__init__.py +40 -0
  225. k3_node/layers/dense/dense_gat_conv.py +149 -0
  226. k3_node/layers/dense/dense_gcn_conv.py +117 -0
  227. k3_node/layers/dense/dense_gin_conv.py +88 -0
  228. k3_node/layers/dense/dense_graph_conv.py +95 -0
  229. k3_node/layers/dense/dense_sage_conv.py +85 -0
  230. k3_node/layers/dense/diff_pool.py +76 -0
  231. k3_node/layers/dense/dmon_pool.py +223 -0
  232. k3_node/layers/dense/linear.py +327 -0
  233. k3_node/layers/dense/mincut_pool.py +92 -0
  234. k3_node/layers/dense/test_dense.py +377 -0
  235. k3_node/layers/functional/__init__.py +13 -0
  236. k3_node/layers/functional/bro.py +49 -0
  237. k3_node/layers/functional/edge_dropout.py +55 -0
  238. k3_node/layers/functional/gini.py +44 -0
  239. k3_node/layers/functional/test_functional.py +34 -0
  240. k3_node/layers/kge/__init__.py +17 -0
  241. k3_node/layers/kge/base.py +255 -0
  242. k3_node/layers/kge/complex.py +98 -0
  243. k3_node/layers/kge/distmult.py +79 -0
  244. k3_node/layers/kge/loader.py +50 -0
  245. k3_node/layers/kge/rotate.py +103 -0
  246. k3_node/layers/kge/test_kge.py +76 -0
  247. k3_node/layers/kge/transe.py +96 -0
  248. k3_node/layers/norm/__init__.py +23 -0
  249. k3_node/layers/norm/batch_norm.py +328 -0
  250. k3_node/layers/norm/diff_group_norm.py +141 -0
  251. k3_node/layers/norm/graph_norm.py +105 -0
  252. k3_node/layers/norm/graph_size_norm.py +57 -0
  253. k3_node/layers/norm/instance_norm.py +163 -0
  254. k3_node/layers/norm/layer_norm.py +245 -0
  255. k3_node/layers/norm/mean_subtraction_norm.py +57 -0
  256. k3_node/layers/norm/msg_norm.py +58 -0
  257. k3_node/layers/norm/pair_norm.py +94 -0
  258. k3_node/layers/norm/test_norm.py +275 -0
  259. k3_node/layers/pool/__init__.py +83 -0
  260. k3_node/layers/pool/approx_knn.py +101 -0
  261. k3_node/layers/pool/asap.py +173 -0
  262. k3_node/layers/pool/avg_pool.py +165 -0
  263. k3_node/layers/pool/cluster_pool.py +168 -0
  264. k3_node/layers/pool/connect/__init__.py +10 -0
  265. k3_node/layers/pool/connect/base.py +103 -0
  266. k3_node/layers/pool/connect/filter_edges.py +113 -0
  267. k3_node/layers/pool/consecutive.py +30 -0
  268. k3_node/layers/pool/decimation.py +48 -0
  269. k3_node/layers/pool/edge_pool.py +189 -0
  270. k3_node/layers/pool/glob.py +139 -0
  271. k3_node/layers/pool/graclus.py +66 -0
  272. k3_node/layers/pool/knn.py +253 -0
  273. k3_node/layers/pool/max_pool.py +159 -0
  274. k3_node/layers/pool/mem_pool.py +145 -0
  275. k3_node/layers/pool/pan_pool.py +144 -0
  276. k3_node/layers/pool/point_cloud.py +212 -0
  277. k3_node/layers/pool/pool.py +119 -0
  278. k3_node/layers/pool/sag_pool.py +174 -0
  279. k3_node/layers/pool/select/__init__.py +10 -0
  280. k3_node/layers/pool/select/base.py +112 -0
  281. k3_node/layers/pool/select/topk.py +206 -0
  282. k3_node/layers/pool/test_pool.py +456 -0
  283. k3_node/layers/pool/topk_pool.py +103 -0
  284. k3_node/layers/pool/voxel_grid.py +70 -0
  285. k3_node/layers/unpool/__init__.py +9 -0
  286. k3_node/layers/unpool/knn_interpolate.py +57 -0
  287. k3_node/layers/unpool/test_unpool.py +31 -0
  288. k3_node/loader/__init__.py +62 -0
  289. k3_node/loader/base.py +69 -0
  290. k3_node/loader/cache.py +68 -0
  291. k3_node/loader/cluster.py +127 -0
  292. k3_node/loader/data_list_loader.py +45 -0
  293. k3_node/loader/dataloader.py +117 -0
  294. k3_node/loader/dense_data_loader.py +62 -0
  295. k3_node/loader/dynamic_batch_sampler.py +93 -0
  296. k3_node/loader/graph_saint.py +188 -0
  297. k3_node/loader/hgt_loader.py +90 -0
  298. k3_node/loader/imbalanced_sampler.py +87 -0
  299. k3_node/loader/keras_dataset.py +334 -0
  300. k3_node/loader/link_loader.py +179 -0
  301. k3_node/loader/link_neighbor_loader.py +202 -0
  302. k3_node/loader/mixin.py +190 -0
  303. k3_node/loader/neighbor_loader.py +159 -0
  304. k3_node/loader/neighbor_sampler.py +167 -0
  305. k3_node/loader/node_loader.py +185 -0
  306. k3_node/loader/prefetch.py +115 -0
  307. k3_node/loader/random_node_loader.py +89 -0
  308. k3_node/loader/sampler_utils.py +499 -0
  309. k3_node/loader/shadow.py +115 -0
  310. k3_node/loader/temporal_dataloader.py +98 -0
  311. k3_node/loader/test_dataloader.py +113 -0
  312. k3_node/loader/test_keras_dataset.py +221 -0
  313. k3_node/loader/test_neighbor_loader.py +122 -0
  314. k3_node/loader/test_sampler_utils.py +82 -0
  315. k3_node/loader/test_samplers.py +96 -0
  316. k3_node/loader/test_subgraph_loaders.py +89 -0
  317. k3_node/loader/utils.py +232 -0
  318. k3_node/loader/zip_loader.py +88 -0
  319. k3_node/metrics.py +94 -0
  320. k3_node/models/__init__.py +424 -0
  321. k3_node/models/attentive_fp.py +232 -0
  322. k3_node/models/attract_repel.py +108 -0
  323. k3_node/models/autoencoder.py +318 -0
  324. k3_node/models/basic_gnn.py +443 -0
  325. k3_node/models/bio/__init__.py +4 -0
  326. k3_node/models/captum.py +52 -0
  327. k3_node/models/chemistry/__init__.py +4 -0
  328. k3_node/models/correct_and_smooth.py +146 -0
  329. k3_node/models/deep_graph_infomax.py +113 -0
  330. k3_node/models/deepgcn.py +121 -0
  331. k3_node/models/dimenet.py +737 -0
  332. k3_node/models/dimenet_utils.py +153 -0
  333. k3_node/models/gnnff.py +263 -0
  334. k3_node/models/gps_model.py +1122 -0
  335. k3_node/models/gpse.py +638 -0
  336. k3_node/models/graph_unet.py +199 -0
  337. k3_node/models/graphmae2.py +954 -0
  338. k3_node/models/graphormer.py +1258 -0
  339. k3_node/models/graphormer_3d.py +868 -0
  340. k3_node/models/grover.py +1066 -0
  341. k3_node/models/jumping_knowledge.py +200 -0
  342. k3_node/models/label_prop.py +110 -0
  343. k3_node/models/lightgcn.py +171 -0
  344. k3_node/models/linkx.py +181 -0
  345. k3_node/models/lpformer.py +404 -0
  346. k3_node/models/mask_label.py +114 -0
  347. k3_node/models/materials/__init__.py +33 -0
  348. k3_node/models/meta.py +133 -0
  349. k3_node/models/metapath2vec.py +234 -0
  350. k3_node/models/mlp.py +264 -0
  351. k3_node/models/mole_bert.py +379 -0
  352. k3_node/models/neural_fingerprint.py +95 -0
  353. k3_node/models/node2vec.py +213 -0
  354. k3_node/models/pmlp.py +157 -0
  355. k3_node/models/polynormer.py +229 -0
  356. k3_node/models/rect.py +93 -0
  357. k3_node/models/renet.py +221 -0
  358. k3_node/models/rev_gnn.py +128 -0
  359. k3_node/models/schnet.py +484 -0
  360. k3_node/models/sgformer.py +195 -0
  361. k3_node/models/signed_gcn.py +185 -0
  362. k3_node/models/test_attentive_fp.py +32 -0
  363. k3_node/models/test_attract_repel.py +33 -0
  364. k3_node/models/test_autoencoder.py +119 -0
  365. k3_node/models/test_basic_gnn.py +102 -0
  366. k3_node/models/test_correct_and_smooth.py +40 -0
  367. k3_node/models/test_deep_graph_infomax.py +68 -0
  368. k3_node/models/test_deepgcn.py +21 -0
  369. k3_node/models/test_dimenet.py +86 -0
  370. k3_node/models/test_domain_apis.py +138 -0
  371. k3_node/models/test_gnnff.py +24 -0
  372. k3_node/models/test_gps_model.py +271 -0
  373. k3_node/models/test_gpse.py +34 -0
  374. k3_node/models/test_graph_unet.py +26 -0
  375. k3_node/models/test_graphmae2.py +226 -0
  376. k3_node/models/test_graphormer.py +233 -0
  377. k3_node/models/test_graphormer3d.py +163 -0
  378. k3_node/models/test_grover.py +287 -0
  379. k3_node/models/test_jumping_knowledge.py +129 -0
  380. k3_node/models/test_label_prop.py +37 -0
  381. k3_node/models/test_lightgcn.py +38 -0
  382. k3_node/models/test_linkx.py +31 -0
  383. k3_node/models/test_lpformer.py +22 -0
  384. k3_node/models/test_mask_label.py +90 -0
  385. k3_node/models/test_meta.py +159 -0
  386. k3_node/models/test_metapath2vec.py +45 -0
  387. k3_node/models/test_mlp.py +62 -0
  388. k3_node/models/test_mole_bert.py +164 -0
  389. k3_node/models/test_neural_fingerprint.py +13 -0
  390. k3_node/models/test_node2vec.py +57 -0
  391. k3_node/models/test_pmlp.py +81 -0
  392. k3_node/models/test_polynormer.py +104 -0
  393. k3_node/models/test_rect.py +23 -0
  394. k3_node/models/test_renet.py +32 -0
  395. k3_node/models/test_rev_gnn.py +24 -0
  396. k3_node/models/test_schnet.py +43 -0
  397. k3_node/models/test_sgformer.py +48 -0
  398. k3_node/models/test_signed_gcn.py +28 -0
  399. k3_node/models/test_tgn.py +77 -0
  400. k3_node/models/test_unimol.py +179 -0
  401. k3_node/models/test_unimol2.py +114 -0
  402. k3_node/models/test_unimol_plus.py +131 -0
  403. k3_node/models/test_visnet.py +44 -0
  404. k3_node/models/tgn.py +382 -0
  405. k3_node/models/unimol.py +1156 -0
  406. k3_node/models/unimol2.py +616 -0
  407. k3_node/models/unimol_docking_v2.py +301 -0
  408. k3_node/models/unimol_plus.py +456 -0
  409. k3_node/models/utils.py +97 -0
  410. k3_node/models/visnet.py +759 -0
  411. k3_node/ops/__init__.py +4 -0
  412. k3_node/ops/conv.py +56 -0
  413. k3_node/ops/creation.py +43 -0
  414. k3_node/ops/graph.py +27 -0
  415. k3_node/ops/host.py +41 -0
  416. k3_node/ops/matmul.py +49 -0
  417. k3_node/ops/numpy.py +24 -0
  418. k3_node/ops/segment.py +54 -0
  419. k3_node/ops/sparse.py +51 -0
  420. k3_node/rag/__init__.py +49 -0
  421. k3_node/rag/encoders.py +312 -0
  422. k3_node/rag/pipeline.py +192 -0
  423. k3_node/rag/projector.py +184 -0
  424. k3_node/rag/subgraph.py +270 -0
  425. k3_node/rag/test_rag.py +347 -0
  426. k3_node/rag/verbalizer.py +162 -0
  427. k3_node/tasks/__init__.py +19 -0
  428. k3_node/tasks/backbone_resolver.py +125 -0
  429. k3_node/tasks/base.py +67 -0
  430. k3_node/tasks/graph_classification.py +270 -0
  431. k3_node/tasks/graph_regression.py +228 -0
  432. k3_node/tasks/link_prediction.py +306 -0
  433. k3_node/tasks/node_classification.py +194 -0
  434. k3_node/tasks/node_regression.py +138 -0
  435. k3_node/tasks/test_tasks.py +319 -0
  436. k3_node/test_docstring_examples.py +106 -0
  437. k3_node/test_training_forwarding.py +116 -0
  438. k3_node/training.py +115 -0
  439. k3_node/transforms/__init__.py +166 -0
  440. k3_node/transforms/base_transform.py +32 -0
  441. k3_node/transforms/compose.py +58 -0
  442. k3_node/transforms/general.py +676 -0
  443. k3_node/transforms/graph.py +1070 -0
  444. k3_node/transforms/spatial.py +797 -0
  445. k3_node/transforms/test_random_link_split.py +45 -0
  446. k3_node/transforms/test_spatial_transforms.py +65 -0
  447. k3_node/transforms/test_transforms.py +253 -0
  448. k3_node/transforms/utils.py +102 -0
  449. k3_node/utils/__init__.py +5 -0
  450. k3_node/utils/backend_import.py +12 -0
  451. k3_node/utils/graph.py +286 -0
  452. k3_node/utils/keras.py +94 -0
  453. k3_node/utils/random.py +103 -0
  454. k3_node/utils/smiles.py +235 -0
  455. k3_node-1.0.0.dist-info/METADATA +284 -0
  456. k3_node-1.0.0.dist-info/RECORD +459 -0
  457. k3_node-1.0.0.dist-info/WHEEL +5 -0
  458. k3_node-1.0.0.dist-info/licenses/LICENSE +21 -0
  459. k3_node-1.0.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,115 @@
1
+ from typing import Any, Type, TypeVar
2
+
3
+ from keras import ops
4
+
5
+ from k3_node.data.data import BaseData, Data
6
+ from k3_node.data.hetero_data import HeteroData
7
+ from k3_node.data.storage import BaseStorage, get_shape, is_tensor_like
8
+
9
+ T = TypeVar("T")
10
+
11
+
12
+ def narrow(tensor, dim: int, start: int, length: int):
13
+ shape = list(get_shape(tensor))
14
+ if dim < 0:
15
+ dim = len(shape) + dim
16
+ slices = [slice(None)] * len(shape)
17
+ slices[dim] = slice(start, start + length)
18
+ return tensor[tuple(slices)]
19
+
20
+
21
+ def separate(
22
+ cls: Type[T],
23
+ batch: Any,
24
+ idx: int,
25
+ slice_dict: Any,
26
+ inc_dict: Any = None,
27
+ decrement: bool = True,
28
+ ) -> T:
29
+ is_hetero = isinstance(batch, HeteroData)
30
+ data = HeteroData() if is_hetero else Data()
31
+
32
+ if not is_hetero:
33
+ batch_store = batch._store
34
+ data_store = data._store
35
+
36
+ for attr in slice_dict.keys():
37
+ if attr not in batch_store:
38
+ continue
39
+ slices = slice_dict[attr]
40
+ incs = inc_dict[attr] if (decrement and inc_dict is not None and attr in inc_dict) else None
41
+
42
+ val = batch_store[attr]
43
+ if is_tensor_like(val):
44
+ start = int(slices[idx])
45
+ end = int(slices[idx + 1])
46
+ length = end - start
47
+ cat_dim = batch.__cat_dim__(attr, val, batch_store)
48
+ sub_val = narrow(val, cat_dim, start, length)
49
+ if decrement and incs is not None and idx < len(incs):
50
+ inc_val = incs[idx]
51
+ if inc_val != 0:
52
+ inc_t = ops.convert_to_tensor(inc_val, dtype=sub_val.dtype)
53
+ sub_val = sub_val - inc_t
54
+ data_store[attr] = sub_val
55
+ elif isinstance(val, (list, tuple)) and len(val) > idx:
56
+ data_store[attr] = val[idx]
57
+ else:
58
+ data_store[attr] = val
59
+
60
+ if hasattr(batch_store, "_num_nodes") and idx < len(batch_store._num_nodes):
61
+ data_store.num_nodes = batch_store._num_nodes[idx]
62
+
63
+ else:
64
+ # Heterogeneous separate
65
+ for node_type in batch.node_types:
66
+ batch_store = batch[node_type]
67
+ data_store = data[node_type]
68
+ store_slice_dict = slice_dict.get(node_type, {})
69
+ store_inc_dict = inc_dict.get(node_type, {}) if decrement and inc_dict else {}
70
+
71
+ for attr in store_slice_dict.keys():
72
+ if attr not in batch_store:
73
+ continue
74
+ slices = store_slice_dict[attr]
75
+ val = batch_store[attr]
76
+ if is_tensor_like(val):
77
+ start = int(slices[idx])
78
+ end = int(slices[idx + 1])
79
+ length = end - start
80
+ cat_dim = batch.__cat_dim__(attr, val, batch_store)
81
+ sub_val = narrow(val, cat_dim, start, length)
82
+ data_store[attr] = sub_val
83
+ elif isinstance(val, (list, tuple)) and len(val) > idx:
84
+ data_store[attr] = val[idx]
85
+
86
+ for edge_type in batch.edge_types:
87
+ batch_store = batch[edge_type]
88
+ data_store = data[edge_type]
89
+ store_slice_dict = slice_dict.get(edge_type, {})
90
+ store_inc_dict = inc_dict.get(edge_type, {}) if decrement and inc_dict else {}
91
+
92
+ for attr in store_slice_dict.keys():
93
+ if attr not in batch_store:
94
+ continue
95
+ slices = store_slice_dict[attr]
96
+ incs = store_inc_dict.get(attr) if decrement else None
97
+ val = batch_store[attr]
98
+ if is_tensor_like(val):
99
+ start = int(slices[idx])
100
+ end = int(slices[idx + 1])
101
+ length = end - start
102
+ cat_dim = batch.__cat_dim__(attr, val, batch_store)
103
+ sub_val = narrow(val, cat_dim, start, length)
104
+ if decrement and incs is not None and idx < len(incs):
105
+ inc_val = incs[idx]
106
+ if is_tensor_like(inc_val) or (isinstance(inc_val, np.ndarray) and np.any(inc_val != 0)):
107
+ if hasattr(inc_val, "ndim") and inc_val.ndim == 1:
108
+ inc_val = inc_val[:, None]
109
+ inc_t = ops.convert_to_tensor(inc_val, dtype=sub_val.dtype)
110
+ sub_val = sub_val - inc_t
111
+ data_store[attr] = sub_val
112
+ elif isinstance(val, (list, tuple)) and len(val) > idx:
113
+ data_store[attr] = val[idx]
114
+
115
+ return data
@@ -0,0 +1,593 @@
1
+ import copy
2
+ import weakref
3
+ from collections import defaultdict
4
+ from collections.abc import Mapping, MutableMapping, Sequence
5
+ from enum import Enum
6
+ from typing import Any, Callable, Dict, Iterator, List, Optional, Set, Tuple, Union
7
+
8
+ import numpy as np
9
+ from keras import ops
10
+
11
+ from k3_node.data.view import ItemsView, KeysView, ValuesView
12
+ from k3_node.utils.graph import coalesce, contains_isolated_nodes, has_self_loops, is_undirected
13
+
14
+ N_KEYS = {"x", "feat", "pos", "batch", "node_type", "n_id", "tf"}
15
+ E_KEYS = {"edge_index", "edge_weight", "edge_attr", "edge_type", "e_id"}
16
+
17
+
18
+ def is_tensor_like(x: Any) -> bool:
19
+ if isinstance(x, np.ndarray):
20
+ return True
21
+ return hasattr(x, "shape") and hasattr(x, "dtype")
22
+
23
+
24
+ def to_numpy(x: Any) -> Any:
25
+ if x is None:
26
+ return None
27
+ if hasattr(x, "detach"):
28
+ x = x.detach()
29
+ if hasattr(x, "numpy"):
30
+ try:
31
+ return x.numpy()
32
+ except TypeError:
33
+ if hasattr(x, "cpu"):
34
+ return x.cpu().numpy()
35
+ raise
36
+ if hasattr(x, "cpu"):
37
+ x = x.cpu()
38
+ if hasattr(x, "numpy"):
39
+ return x.numpy()
40
+ if hasattr(x, "_numpy"):
41
+ return x._numpy()
42
+ return np.asarray(x)
43
+
44
+
45
+ def get_shape(x: Any) -> Tuple[int, ...]:
46
+ if hasattr(x, "shape"):
47
+ return tuple(int(s) if s is not None else 0 for s in x.shape)
48
+ if isinstance(x, (list, tuple)):
49
+ return (len(x),)
50
+ return ()
51
+
52
+
53
+ def recursive_apply(data: Any, func: Callable) -> Any:
54
+ if is_tensor_like(data):
55
+ return func(data)
56
+ elif isinstance(data, tuple) and hasattr(data, "_fields"):
57
+ return type(data)(*(recursive_apply(d, func) for d in data))
58
+ elif isinstance(data, Sequence) and not isinstance(data, str):
59
+ return [recursive_apply(d, func) for d in data]
60
+ elif isinstance(data, Mapping):
61
+ return {key: recursive_apply(data[key], func) for key in data}
62
+ else:
63
+ try:
64
+ return func(data)
65
+ except Exception:
66
+ return data
67
+
68
+
69
+ def recursive_apply_(data: Any, func: Callable) -> Any:
70
+ if is_tensor_like(data):
71
+ try:
72
+ func(data)
73
+ except Exception:
74
+ pass
75
+ elif isinstance(data, tuple) and hasattr(data, "_fields"):
76
+ for value in data:
77
+ recursive_apply_(value, func)
78
+ elif isinstance(data, Sequence) and not isinstance(data, str):
79
+ for value in data:
80
+ recursive_apply_(value, func)
81
+ elif isinstance(data, Mapping):
82
+ for value in data.values():
83
+ recursive_apply_(value, func)
84
+ else:
85
+ try:
86
+ func(data)
87
+ except Exception:
88
+ pass
89
+
90
+
91
+ class AttrType(Enum):
92
+ NODE = "NODE"
93
+ EDGE = "EDGE"
94
+ OTHER = "OTHER"
95
+
96
+
97
+ class BaseStorage(MutableMapping):
98
+ def __init__(self, _mapping: Optional[Dict[str, Any]] = None, **kwargs: Any) -> None:
99
+ super().__init__()
100
+ self._mapping: Dict[str, Any] = {}
101
+ for key, value in (_mapping or {}).items():
102
+ setattr(self, key, value)
103
+ for key, value in kwargs.items():
104
+ setattr(self, key, value)
105
+
106
+ @property
107
+ def _key(self) -> Any:
108
+ return None
109
+
110
+ def _pop_cache(self, key: str) -> None:
111
+ for cache in getattr(self, "_cached_attr", {}).values():
112
+ cache.discard(key)
113
+
114
+ def __len__(self) -> int:
115
+ return len(self._mapping)
116
+
117
+ def __getattr__(self, key: str) -> Any:
118
+ if key == "_mapping":
119
+ self._mapping = {}
120
+ return self._mapping
121
+ try:
122
+ return self[key]
123
+ except KeyError:
124
+ raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'") from None
125
+
126
+ def __setattr__(self, key: str, value: Any) -> None:
127
+ propobj = getattr(self.__class__, key, None)
128
+ if propobj is not None and getattr(propobj, "fset", None) is not None:
129
+ propobj.fset(self, value)
130
+ elif key == "_parent":
131
+ self.__dict__[key] = weakref.ref(value) if value is not None else None
132
+ elif key[:1] == "_":
133
+ self.__dict__[key] = value
134
+ else:
135
+ self[key] = value
136
+
137
+ def __delattr__(self, key: str) -> None:
138
+ if key[:1] == "_":
139
+ if key in self.__dict__:
140
+ del self.__dict__[key]
141
+ else:
142
+ del self[key]
143
+
144
+ def __getitem__(self, key: str) -> Any:
145
+ return self._mapping[key]
146
+
147
+ def __setitem__(self, key: str, value: Any) -> None:
148
+ self._pop_cache(key)
149
+ if value is None and key in self._mapping:
150
+ del self._mapping[key]
151
+ elif value is not None:
152
+ self._mapping[key] = value
153
+
154
+ def __delitem__(self, key: str) -> None:
155
+ if key in self._mapping:
156
+ self._pop_cache(key)
157
+ del self._mapping[key]
158
+
159
+ def __iter__(self) -> Iterator[Any]:
160
+ return iter(self._mapping)
161
+
162
+ def __copy__(self):
163
+ out = self.__class__.__new__(self.__class__)
164
+ for key, value in self.__dict__.items():
165
+ if key != "_cached_attr":
166
+ out.__dict__[key] = value
167
+ out._mapping = copy.copy(out._mapping)
168
+ return out
169
+
170
+ def __deepcopy__(self, memo=None):
171
+ out = self.__class__.__new__(self.__class__)
172
+ for key, value in self.__dict__.items():
173
+ if key == "_parent":
174
+ out.__dict__[key] = self.__dict__[key]
175
+ elif key != "_cached_attr":
176
+ out.__dict__[key] = copy.deepcopy(value, memo)
177
+ out._mapping = copy.deepcopy(out._mapping, memo)
178
+ return out
179
+
180
+ def __getstate__(self) -> Dict[str, Any]:
181
+ out = self.__dict__.copy()
182
+ _parent = out.get("_parent", None)
183
+ if _parent is not None:
184
+ out["_parent"] = _parent()
185
+ return out
186
+
187
+ def __setstate__(self, mapping: Dict[str, Any]) -> None:
188
+ for key, value in mapping.items():
189
+ self.__dict__[key] = value
190
+ _parent = self.__dict__.get("_parent", None)
191
+ if _parent is not None:
192
+ self.__dict__["_parent"] = weakref.ref(_parent)
193
+
194
+ def __repr__(self) -> str:
195
+ return repr(self._mapping)
196
+
197
+ def _parent(self):
198
+ parent_ref = self.__dict__.get("_parent", None)
199
+ return parent_ref() if parent_ref is not None else None
200
+
201
+ def keys(self, *args: str) -> KeysView:
202
+ return KeysView(self._mapping, *args)
203
+
204
+ def values(self, *args: str) -> ValuesView:
205
+ return ValuesView(self._mapping, *args)
206
+
207
+ def items(self, *args: str) -> ItemsView:
208
+ return ItemsView(self._mapping, *args)
209
+
210
+ def to_dict(self) -> Dict[str, Any]:
211
+ return copy.copy(self._mapping)
212
+
213
+ def apply_(self, func: Callable, *args: str):
214
+ for value in self.values(*args):
215
+ recursive_apply_(value, func)
216
+ return self
217
+
218
+ def apply(self, func: Callable, *args: str):
219
+ for key, value in self.items(*args):
220
+ self[key] = recursive_apply(value, func)
221
+ return self
222
+
223
+ def to(self, *args, **kwargs):
224
+ def _to(x):
225
+ if hasattr(x, "to"):
226
+ return x.to(*args, **kwargs)
227
+ return x
228
+
229
+ return self.apply(_to)
230
+
231
+ def to_backend(self, backend: Optional[str] = None):
232
+ for key in list(self.keys()):
233
+ val = self[key]
234
+ if is_tensor_like(val):
235
+ v_np = to_numpy(val)
236
+ if v_np.dtype == np.bool_:
237
+ self[key] = ops.convert_to_tensor(v_np, dtype="bool")
238
+ elif np.issubdtype(v_np.dtype, np.integer):
239
+ self[key] = ops.convert_to_tensor(v_np, dtype="int64")
240
+ elif np.issubdtype(v_np.dtype, np.floating):
241
+ self[key] = ops.convert_to_tensor(v_np, dtype="float32")
242
+ else:
243
+ self[key] = ops.convert_to_tensor(v_np)
244
+ return self
245
+
246
+ def cpu(self):
247
+ return self.to("cpu")
248
+
249
+ def cuda(self):
250
+ return self.to("cuda")
251
+
252
+ def requires_grad_(self, *args: str):
253
+ def _req(x):
254
+ if hasattr(x, "requires_grad_"):
255
+ x.requires_grad_()
256
+ elif hasattr(x, "requires_grad"):
257
+ x.requires_grad = True
258
+
259
+ return self.apply_(_req, *args)
260
+
261
+ def contiguous(self, *args: str):
262
+ def _cont(x):
263
+ if hasattr(x, "contiguous"):
264
+ return x.contiguous()
265
+ return x
266
+
267
+ return self.apply(_cont, *args)
268
+
269
+
270
+ class NodeStorage(BaseStorage):
271
+ @property
272
+ def _key(self) -> Any:
273
+ return self.__dict__.get("_key", None)
274
+
275
+ @property
276
+ def num_nodes(self) -> int:
277
+ if "num_nodes" in self:
278
+ return int(self["num_nodes"])
279
+ parent = self._parent()
280
+ for key, value in self.items():
281
+ if is_tensor_like(value) and key in N_KEYS:
282
+ cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else 0
283
+ return get_shape(value)[cat_dim]
284
+ for key, value in self.items():
285
+ if is_tensor_like(value) and "node" in key:
286
+ cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else 0
287
+ return get_shape(value)[cat_dim]
288
+ edge_index = self.get("edge_index")
289
+ if is_tensor_like(edge_index) and get_shape(edge_index)[-1] > 0:
290
+ # As in PyG: without node-level attributes, infer the count from the edges
291
+ return int(np.asarray(to_numpy(edge_index)).max()) + 1
292
+ return 0
293
+
294
+ @property
295
+ def num_node_features(self) -> int:
296
+ x = self.get("x")
297
+ if x is not None and is_tensor_like(x):
298
+ shape = get_shape(x)
299
+ return 1 if len(shape) == 1 else shape[-1]
300
+ return 0
301
+
302
+ @property
303
+ def num_features(self) -> int:
304
+ return self.num_node_features
305
+
306
+ def is_node_attr(self, key: str) -> bool:
307
+ if "_cached_attr" not in self.__dict__:
308
+ self._cached_attr: Dict[AttrType, Set[str]] = defaultdict(set)
309
+
310
+ if key in self._cached_attr[AttrType.NODE]:
311
+ return True
312
+ if key in self._cached_attr[AttrType.OTHER]:
313
+ return False
314
+
315
+ value = self.get(key)
316
+ if value is None:
317
+ return False
318
+
319
+ if isinstance(value, (list, tuple)) and len(value) == self.num_nodes:
320
+ self._cached_attr[AttrType.NODE].add(key)
321
+ return True
322
+
323
+ if not is_tensor_like(value):
324
+ self._cached_attr[AttrType.OTHER].add(key)
325
+ return False
326
+
327
+ shape = get_shape(value)
328
+ if len(shape) == 0:
329
+ self._cached_attr[AttrType.OTHER].add(key)
330
+ return False
331
+
332
+ parent = self._parent()
333
+ cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else 0
334
+ if shape[cat_dim] != self.num_nodes:
335
+ self._cached_attr[AttrType.OTHER].add(key)
336
+ return False
337
+
338
+ self._cached_attr[AttrType.NODE].add(key)
339
+ return True
340
+
341
+ def is_edge_attr(self, key: str) -> bool:
342
+ return False
343
+
344
+ def node_attrs(self) -> List[str]:
345
+ return [key for key in self.keys() if self.is_node_attr(key)]
346
+
347
+
348
+ class EdgeStorage(BaseStorage):
349
+ @property
350
+ def _key(self) -> Any:
351
+ return self.__dict__.get("_key", None)
352
+
353
+ @property
354
+ def edge_index(self):
355
+ if "edge_index" in self:
356
+ return self["edge_index"]
357
+ raise AttributeError(f"'{self.__class__.__name__}' object has no attribute 'edge_index'")
358
+
359
+ @edge_index.setter
360
+ def edge_index(self, edge_index) -> None:
361
+ self["edge_index"] = edge_index
362
+
363
+ @property
364
+ def num_edges(self) -> int:
365
+ if "num_edges" in self:
366
+ return int(self["num_edges"])
367
+ parent = self._parent()
368
+ for key, value in self.items():
369
+ if is_tensor_like(value) and key in E_KEYS:
370
+ cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else -1
371
+ return get_shape(value)[cat_dim]
372
+ for key, value in self.items():
373
+ if is_tensor_like(value) and "edge" in key:
374
+ cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else -1
375
+ return get_shape(value)[cat_dim]
376
+ return 0
377
+
378
+ @property
379
+ def num_edge_features(self) -> int:
380
+ edge_attr = self.get("edge_attr")
381
+ if edge_attr is not None and is_tensor_like(edge_attr):
382
+ shape = get_shape(edge_attr)
383
+ return 1 if len(shape) == 1 else shape[-1]
384
+ return 0
385
+
386
+ @property
387
+ def num_features(self) -> int:
388
+ return self.num_edge_features
389
+
390
+ def size(self, dim: Optional[int] = None) -> Union[Tuple[Optional[int], Optional[int]], Optional[int]]:
391
+ parent = self._parent()
392
+ if self._key is None or parent is None:
393
+ num = self.num_edges
394
+ res = (num, num)
395
+ return res if dim is None else res[dim]
396
+ size = (parent[self._key[0]].num_nodes, parent[self._key[-1]].num_nodes)
397
+ return size if dim is None else size[dim]
398
+
399
+ def is_node_attr(self, key: str) -> bool:
400
+ return False
401
+
402
+ def is_edge_attr(self, key: str) -> bool:
403
+ if "_cached_attr" not in self.__dict__:
404
+ self._cached_attr: Dict[AttrType, Set[str]] = defaultdict(set)
405
+
406
+ if key in self._cached_attr[AttrType.EDGE]:
407
+ return True
408
+ if key in self._cached_attr[AttrType.OTHER]:
409
+ return False
410
+
411
+ value = self.get(key)
412
+ if value is None:
413
+ return False
414
+
415
+ if isinstance(value, (list, tuple)) and len(value) == self.num_edges:
416
+ self._cached_attr[AttrType.EDGE].add(key)
417
+ return True
418
+
419
+ if not is_tensor_like(value):
420
+ self._cached_attr[AttrType.OTHER].add(key)
421
+ return False
422
+
423
+ shape = get_shape(value)
424
+ if len(shape) == 0:
425
+ self._cached_attr[AttrType.OTHER].add(key)
426
+ return False
427
+
428
+ parent = self._parent()
429
+ cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else -1
430
+ if shape[cat_dim] != self.num_edges:
431
+ self._cached_attr[AttrType.OTHER].add(key)
432
+ return False
433
+
434
+ self._cached_attr[AttrType.EDGE].add(key)
435
+ return True
436
+
437
+ def edge_attrs(self) -> List[str]:
438
+ return [key for key in self.keys() if self.is_edge_attr(key)]
439
+
440
+ def is_coalesced(self) -> bool:
441
+ if "edge_index" in self:
442
+ edge_index = self.edge_index
443
+ new_edge_index, _ = coalesce(edge_index)
444
+ orig_np = ops.convert_to_numpy(edge_index)
445
+ new_np = ops.convert_to_numpy(new_edge_index)
446
+ return orig_np.shape == new_np.shape and np.array_equal(orig_np, new_np)
447
+ return True
448
+
449
+ def coalesce(self, reduce: str = "add"):
450
+ if "edge_index" in self:
451
+ self.edge_index, self.edge_attr = coalesce(
452
+ self.edge_index,
453
+ edge_attr=self.get("edge_attr"),
454
+ reduce=reduce,
455
+ )
456
+ return self
457
+
458
+ def has_self_loops(self) -> bool:
459
+ if self.is_bipartite() or "edge_index" not in self:
460
+ return False
461
+ return has_self_loops(self.edge_index)
462
+
463
+ def has_isolated_nodes(self) -> bool:
464
+ if "edge_index" not in self:
465
+ return False
466
+ parent = self._parent()
467
+ num_nodes = parent[self._key[-1]].num_nodes if parent and self._key else None
468
+ return contains_isolated_nodes(self.edge_index, num_nodes=num_nodes)
469
+
470
+ def is_undirected(self) -> bool:
471
+ if self.is_bipartite() or "edge_index" not in self:
472
+ return False
473
+ return is_undirected(self.edge_index, edge_attr=self.get("edge_attr"))
474
+
475
+ def is_directed(self) -> bool:
476
+ return not self.is_undirected()
477
+
478
+ def is_bipartite(self) -> bool:
479
+ return self._key is not None and isinstance(self._key, tuple) and self._key[0] != self._key[-1]
480
+
481
+
482
+ class GlobalStorage(NodeStorage, EdgeStorage):
483
+ @property
484
+ def _key(self) -> Any:
485
+ return None
486
+
487
+ @property
488
+ def num_features(self) -> int:
489
+ return self.num_node_features
490
+
491
+ def size(self, dim: Optional[int] = None) -> Union[Tuple[Optional[int], Optional[int]], Optional[int]]:
492
+ size = (self.num_nodes, self.num_nodes)
493
+ return size if dim is None else size[dim]
494
+
495
+ def is_node_attr(self, key: str) -> bool:
496
+ if "_cached_attr" not in self.__dict__:
497
+ self._cached_attr: Dict[AttrType, Set[str]] = defaultdict(set)
498
+
499
+ if key in self._cached_attr[AttrType.NODE]:
500
+ return True
501
+ if key in self._cached_attr[AttrType.EDGE] or key in self._cached_attr[AttrType.OTHER]:
502
+ return False
503
+
504
+ value = self.get(key)
505
+ if value is None:
506
+ return False
507
+
508
+ if isinstance(value, (list, tuple)) and len(value) == self.num_nodes:
509
+ self._cached_attr[AttrType.NODE].add(key)
510
+ return True
511
+
512
+ if not is_tensor_like(value):
513
+ return False
514
+
515
+ shape = get_shape(value)
516
+ if len(shape) == 0:
517
+ self._cached_attr[AttrType.OTHER].add(key)
518
+ return False
519
+
520
+ parent = self._parent()
521
+ cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else 0
522
+ if not isinstance(cat_dim, int):
523
+ return False
524
+
525
+ num_nodes, num_edges = self.num_nodes, self.num_edges
526
+
527
+ if shape[cat_dim] != num_nodes:
528
+ if shape[cat_dim] == num_edges:
529
+ self._cached_attr[AttrType.EDGE].add(key)
530
+ else:
531
+ self._cached_attr[AttrType.OTHER].add(key)
532
+ return False
533
+
534
+ if num_nodes != num_edges:
535
+ self._cached_attr[AttrType.NODE].add(key)
536
+ return True
537
+
538
+ if "edge" not in key:
539
+ self._cached_attr[AttrType.NODE].add(key)
540
+ return True
541
+ else:
542
+ self._cached_attr[AttrType.EDGE].add(key)
543
+ return False
544
+
545
+ def is_edge_attr(self, key: str) -> bool:
546
+ if "_cached_attr" not in self.__dict__:
547
+ self._cached_attr = defaultdict(set)
548
+
549
+ if key in self._cached_attr[AttrType.EDGE]:
550
+ return True
551
+ if key in self._cached_attr[AttrType.NODE] or key in self._cached_attr[AttrType.OTHER]:
552
+ return False
553
+
554
+ value = self.get(key)
555
+ if value is None:
556
+ return False
557
+
558
+ if isinstance(value, (list, tuple)) and len(value) == self.num_edges:
559
+ self._cached_attr[AttrType.EDGE].add(key)
560
+ return True
561
+
562
+ if not is_tensor_like(value):
563
+ return False
564
+
565
+ shape = get_shape(value)
566
+ if len(shape) == 0:
567
+ self._cached_attr[AttrType.OTHER].add(key)
568
+ return False
569
+
570
+ parent = self._parent()
571
+ cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else -1
572
+ if not isinstance(cat_dim, int):
573
+ return False
574
+
575
+ num_nodes, num_edges = self.num_nodes, self.num_edges
576
+
577
+ if shape[cat_dim] != num_edges:
578
+ if shape[cat_dim] == num_nodes:
579
+ self._cached_attr[AttrType.NODE].add(key)
580
+ else:
581
+ self._cached_attr[AttrType.OTHER].add(key)
582
+ return False
583
+
584
+ if num_edges != num_nodes:
585
+ self._cached_attr[AttrType.EDGE].add(key)
586
+ return True
587
+
588
+ if "edge" in key:
589
+ self._cached_attr[AttrType.EDGE].add(key)
590
+ return True
591
+ else:
592
+ self._cached_attr[AttrType.NODE].add(key)
593
+ return False