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,96 @@
1
+ import os
2
+ import os.path as osp
3
+ from itertools import product
4
+ from typing import Callable, List, Optional
5
+ import numpy as np
6
+ from keras import ops
7
+
8
+ from k3_node.data import HeteroData, InMemoryDataset
9
+ from k3_node.io import fs
10
+
11
+
12
+ class IMDB(InMemoryDataset):
13
+ r"""A subset of the Internet Movie Database (IMDB) containing three types of entities:
14
+ movies, actors, and directors.
15
+ """
16
+
17
+ url = "https://www.dropbox.com/s/g0btk9ctr1es39x/IMDB_processed.zip?dl=1"
18
+
19
+ def __init__(
20
+ self,
21
+ root: str,
22
+ transform: Optional[Callable] = None,
23
+ pre_transform: Optional[Callable] = None,
24
+ force_reload: bool = False,
25
+ ):
26
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
27
+ self.load(self.processed_paths[0])
28
+
29
+ @property
30
+ def raw_file_names(self) -> List[str]:
31
+ return [
32
+ "adjM.npz",
33
+ "features_0.npz",
34
+ "features_1.npz",
35
+ "features_2.npz",
36
+ "labels.npy",
37
+ "train_val_test_idx.npz",
38
+ ]
39
+
40
+ @property
41
+ def processed_file_names(self) -> str:
42
+ return "data.pt"
43
+
44
+ def download(self):
45
+ zip_path = osp.join(self.raw_dir, "IMDB_processed.zip")
46
+ fs.cp(self.url, zip_path, extract=True)
47
+ if osp.exists(zip_path):
48
+ fs.rm(zip_path)
49
+
50
+ def process(self):
51
+ import scipy.sparse as sp
52
+
53
+ data = HeteroData()
54
+ node_types = ["movie", "director", "actor"]
55
+
56
+ for i, node_type in enumerate(node_types):
57
+ x = sp.load_npz(osp.join(self.raw_dir, f"features_{i}.npz"))
58
+ data[node_type].x = ops.convert_to_tensor(
59
+ np.array(x.todense(), dtype=np.float32), dtype="float32"
60
+ )
61
+
62
+ y = np.load(osp.join(self.raw_dir, "labels.npy"))
63
+ data["movie"].y = ops.convert_to_tensor(y.astype(np.int64), dtype="int64")
64
+
65
+ split = np.load(osp.join(self.raw_dir, "train_val_test_idx.npz"))
66
+ for name in ["train", "val", "test"]:
67
+ idx = split[f"{name}_idx"]
68
+ mask = np.zeros(data["movie"].num_nodes, dtype=bool)
69
+ mask[idx] = True
70
+ data["movie"][f"{name}_mask"] = ops.convert_to_tensor(mask, dtype="bool")
71
+
72
+ s = {}
73
+ N_m = data["movie"].num_nodes
74
+ N_d = data["director"].num_nodes
75
+ N_a = data["actor"].num_nodes
76
+ s["movie"] = (0, N_m)
77
+ s["director"] = (N_m, N_m + N_d)
78
+ s["actor"] = (N_m + N_d, N_m + N_d + N_a)
79
+
80
+ A = sp.load_npz(osp.join(self.raw_dir, "adjM.npz"))
81
+ for src, dst in product(node_types, node_types):
82
+ A_sub = A[s[src][0] : s[src][1], s[dst][0] : s[dst][1]].tocoo()
83
+ if A_sub.nnz > 0:
84
+ row = np.array(A_sub.row, dtype=np.int64)
85
+ col = np.array(A_sub.col, dtype=np.int64)
86
+ edge_index = np.stack([row, col], axis=0)
87
+ data[src, dst].edge_index = ops.convert_to_tensor(edge_index, dtype="int64")
88
+
89
+ if self.pre_transform is not None:
90
+ data = self.pre_transform(data)
91
+
92
+ self.save([data], self.processed_paths[0])
93
+
94
+ def __repr__(self) -> str:
95
+ return f"{self.__class__.__name__}()"
96
+
@@ -0,0 +1,56 @@
1
+ import os
2
+ import os.path as osp
3
+
4
+ import numpy as np
5
+
6
+
7
+ class JODIEDataset:
8
+ r"""The temporal interaction datasets of `"JODIE: Predicting Dynamic Embedding Trajectory in
9
+ Temporal Interaction Networks" <https://cs.stanford.edu/~srijan/pubs/jodie-kdd2019.pdf>`_:
10
+ ``"wikipedia"``, ``"reddit"``, ``"mooc"`` and ``"lastfm"``.
11
+
12
+ ``dataset[0]`` is a :class:`~k3_node.data.TemporalData` event stream: user ``src`` interacts
13
+ with item ``dst`` (numbered after the users) at time ``t``, with features ``msg`` and label
14
+ ``y``. MOOC (7,144 nodes, 411,749 events, 4 features) is the smallest download (40 MB).
15
+
16
+ Args:
17
+ root (str): Root directory where the dataset should be saved.
18
+ name (str): The name of the dataset.
19
+ """
20
+
21
+ url = "https://snap.stanford.edu/jodie/{}.csv"
22
+ names = ["wikipedia", "reddit", "mooc", "lastfm"]
23
+
24
+ def __init__(self, root: str, name: str):
25
+ self.name = name.lower()
26
+ assert self.name in self.names
27
+ folder = osp.join(root, self.name)
28
+ cache = osp.join(folder, "processed.npz")
29
+ if not osp.exists(cache):
30
+ os.makedirs(folder, exist_ok=True)
31
+ csv = osp.join(folder, f"{self.name}.csv")
32
+ if not osp.exists(csv):
33
+ from k3_node.data.download import download_url
34
+
35
+ download_url(self.url.format(self.name), folder)
36
+ import pandas as pd
37
+
38
+ df = pd.read_csv(csv, skiprows=1, header=None)
39
+ src = df.iloc[:, 0].values.astype(np.int64)
40
+ dst = df.iloc[:, 1].values.astype(np.int64) + int(src.max()) + 1
41
+ np.savez(cache, src=src, dst=dst, t=df.iloc[:, 2].values.astype(np.int64),
42
+ y=df.iloc[:, 3].values.astype(np.int64), msg=df.iloc[:, 4:].values.astype(np.float32))
43
+ self._arrays = dict(np.load(cache))
44
+
45
+ def __len__(self):
46
+ return 1
47
+
48
+ def __getitem__(self, idx):
49
+ from k3_node.data import TemporalData
50
+
51
+ if idx != 0:
52
+ raise IndexError(idx)
53
+ return TemporalData(**self._arrays)
54
+
55
+ def __repr__(self):
56
+ return f"JODIEDataset({self.name})"
@@ -0,0 +1,56 @@
1
+ from typing import Callable, Optional
2
+ import numpy as np
3
+ from keras import ops
4
+
5
+ from k3_node.data import Data, InMemoryDataset
6
+
7
+
8
+ class KarateClub(InMemoryDataset):
9
+ r"""Zachary's karate club network from the `"An Information Flow Model for
10
+ Conflict and Fission in Small Groups" paper, containing 34 nodes and 156
11
+ undirected edges labeled into 4 community classes.
12
+ """
13
+
14
+ def __init__(self, transform: Optional[Callable] = None):
15
+ super().__init__(None, transform)
16
+
17
+ row = [
18
+ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1,
19
+ 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 4, 4, 4,
20
+ 5, 5, 5, 5, 6, 6, 6, 6, 7, 7, 7, 7, 8, 8, 8, 8, 8, 9, 9, 10, 10,
21
+ 10, 11, 12, 12, 13, 13, 13, 13, 13, 14, 14, 15, 15, 16, 16, 17, 17,
22
+ 18, 18, 19, 19, 19, 20, 20, 21, 21, 22, 22, 23, 23, 23, 23, 23, 24,
23
+ 24, 24, 25, 25, 25, 26, 26, 27, 27, 27, 27, 28, 28, 28, 29, 29, 29,
24
+ 29, 30, 30, 30, 30, 31, 31, 31, 31, 31, 31, 32, 32, 32, 32, 32, 32,
25
+ 32, 32, 32, 32, 32, 32, 33, 33, 33, 33, 33, 33, 33, 33, 33, 33, 33,
26
+ 33, 33, 33, 33, 33, 33
27
+ ]
28
+ col = [
29
+ 1, 2, 3, 4, 5, 6, 7, 8, 10, 11, 12, 13, 17, 19, 21, 31, 0, 2, 3, 7,
30
+ 13, 17, 19, 21, 30, 0, 1, 3, 7, 8, 9, 13, 27, 28, 32, 0, 1, 2, 7,
31
+ 12, 13, 0, 6, 10, 0, 6, 10, 16, 0, 4, 5, 16, 0, 1, 2, 3, 0, 2, 30,
32
+ 32, 33, 2, 33, 0, 4, 5, 0, 0, 3, 0, 1, 2, 3, 33, 32, 33, 32, 33, 5,
33
+ 6, 0, 1, 32, 33, 0, 1, 33, 32, 33, 0, 1, 32, 33, 25, 27, 29, 32,
34
+ 33, 25, 27, 31, 23, 24, 31, 29, 33, 2, 23, 24, 33, 2, 31, 33, 23,
35
+ 26, 32, 33, 1, 8, 32, 33, 0, 24, 25, 28, 32, 33, 2, 8, 14, 15, 18,
36
+ 20, 22, 23, 29, 30, 31, 33, 8, 9, 13, 14, 15, 18, 19, 20, 22, 23,
37
+ 26, 27, 28, 29, 30, 31, 32
38
+ ]
39
+ edge_index = ops.convert_to_tensor(np.array([row, col], dtype=np.int64), dtype="int64")
40
+
41
+ y_np = np.array([
42
+ 1, 1, 1, 1, 3, 3, 3, 1, 0, 1, 3, 1, 1, 1, 0, 0, 3, 1, 0, 1, 0, 1,
43
+ 0, 0, 2, 2, 0, 0, 2, 0, 0, 2, 0, 0
44
+ ], dtype=np.int64)
45
+ y = ops.convert_to_tensor(y_np, dtype="int64")
46
+
47
+ x = ops.convert_to_tensor(np.eye(34, dtype=np.float32), dtype="float32")
48
+
49
+ train_mask_np = np.zeros(34, dtype=bool)
50
+ for i in range(int(np.max(y_np)) + 1):
51
+ train_mask_np[np.where(y_np == i)[0][0]] = True
52
+ train_mask = ops.convert_to_tensor(train_mask_np, dtype="bool")
53
+
54
+ data = Data(x=x, edge_index=edge_index, y=y, train_mask=train_mask)
55
+ self.data, self.slices = self.collate([data])
56
+
@@ -0,0 +1,51 @@
1
+ from typing import Callable, Optional
2
+ import numpy as np
3
+ from keras import ops
4
+
5
+ from k3_node.data import Data, InMemoryDataset
6
+ from k3_node.io import fs
7
+
8
+
9
+ class LastFMAsia(InMemoryDataset):
10
+ r"""The LastFM Asia Network dataset."""
11
+
12
+ url = "https://graphmining.ai/datasets/ptg/lastfm_asia.npz"
13
+
14
+ def __init__(
15
+ self,
16
+ root: str,
17
+ transform: Optional[Callable] = None,
18
+ pre_transform: Optional[Callable] = None,
19
+ force_reload: bool = False,
20
+ ):
21
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
22
+ self.load(self.processed_paths[0])
23
+
24
+ @property
25
+ def raw_file_names(self) -> str:
26
+ return "lastfm_asia.npz"
27
+
28
+ @property
29
+ def processed_file_names(self) -> str:
30
+ return "data.pt"
31
+
32
+ def download(self):
33
+ fs.cp(self.url, self.raw_dir)
34
+
35
+ def process(self):
36
+ data = np.load(self.raw_paths[0], allow_pickle=True)
37
+ x = data["features"].astype(np.float32)
38
+ y = data["target"].astype(np.int64)
39
+ edge_index = data["edges"].astype(np.int64).T
40
+
41
+ data_obj = Data(
42
+ x=ops.convert_to_tensor(x, dtype="float32"),
43
+ y=ops.convert_to_tensor(y, dtype="int64"),
44
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
45
+ )
46
+
47
+ if self.pre_transform is not None:
48
+ data_obj = self.pre_transform(data_obj)
49
+
50
+ self.save([data_obj], self.processed_paths[0])
51
+
@@ -0,0 +1,50 @@
1
+ import os
2
+ from typing import Callable, Optional
3
+
4
+ import numpy as np
5
+
6
+ from k3_node.data import Data, InMemoryDataset
7
+
8
+
9
+ class MeshCorrespondence(InMemoryDataset):
10
+ r"""A small shape correspondence dataset in the spirit of FAUST: randomly deformed copies of one
11
+ :class:`GeometricShapes` mesh (the monkey head by default). Every copy is smoothly bent and
12
+ slightly rotated; the task is to recognize every vertex, so ``y`` holds the vertex indices.
13
+
14
+ The graphs hold ``pos`` and mesh faces ``face``; use :class:`~k3_node.transforms.FaceToEdge`
15
+ to connect the vertices. As in FAUST, there are 80 training and 20 test meshes by default.
16
+
17
+ Args:
18
+ root (str): Directory where GeometricShapes is (or will be) downloaded.
19
+ train (bool, optional): Build the training meshes if ``True``, else the test meshes.
20
+ num_meshes (int, optional): Number of meshes. (default: ``80`` / ``20``)
21
+ shape (str, optional): The GeometricShapes mesh to deform. (default: ``"3d_monkey"``)
22
+ transform (callable, optional): A function applied to each mesh when it is accessed.
23
+ pre_transform (callable, optional): A function applied to each mesh once, when built.
24
+ seed (int, optional): Random seed. (default: ``0``)
25
+ """
26
+
27
+ def __init__(self, root: str, train: bool = True, num_meshes: Optional[int] = None, shape: str = "3d_monkey",
28
+ transform: Optional[Callable] = None, pre_transform: Optional[Callable] = None, seed: int = 0):
29
+ super().__init__(None, transform)
30
+ from k3_node.datasets.geometric_shapes import GeometricShapes
31
+ from k3_node.transforms.spatial import _np
32
+
33
+ shapes = GeometricShapes(root, train=True)
34
+ names = sorted(os.listdir(shapes.raw_dir))
35
+ base = next(shapes[i] for i in range(len(shapes)) if names[int(_np(shapes[i].y)[0])] == shape)
36
+ pos, face = _np(base.pos).astype(np.float64), _np(base.face)
37
+ pos = pos / np.abs(pos).max()
38
+
39
+ rng = np.random.default_rng(seed + (0 if train else 1))
40
+ graphs = []
41
+ for _ in range(num_meshes or (80 if train else 20)):
42
+ w = rng.normal(size=(3, 3)) * 1.5
43
+ bent = pos + 0.12 * np.sin(pos @ w + rng.uniform(0, 2 * np.pi, size=3)) # smooth deformation
44
+ angle = np.deg2rad(rng.uniform(-15, 15))
45
+ c, s = np.cos(angle), np.sin(angle)
46
+ rot = np.array([[c, 0, s], [0, 1, 0], [-s, 0, c]])
47
+ data = Data(pos=(bent @ rot.T).astype(np.float32), face=face.copy(),
48
+ y=np.arange(len(pos), dtype=np.int64))
49
+ graphs.append(pre_transform(data) if pre_transform is not None else data)
50
+ self.data, self.slices = self.collate(graphs)
@@ -0,0 +1,148 @@
1
+ import os
2
+ import os.path as osp
3
+ import re
4
+ import warnings
5
+ from typing import Callable, Dict, Optional, Tuple, Union
6
+
7
+ import numpy as np
8
+ from keras import ops
9
+
10
+ from k3_node.data import InMemoryDataset, download_url, extract_gz
11
+ from k3_node.utils.smiles import from_smiles as default_from_smiles
12
+
13
+
14
+ class MoleculeNet(InMemoryDataset):
15
+ r"""The `MoleculeNet <http://moleculenet.org/datasets-1>`_ benchmark collection
16
+ from the `"MoleculeNet: A Benchmark for Molecular Machine Learning"
17
+ <https://arxiv.org/abs/1703.00564>`_ paper, containing datasets from physical
18
+ chemistry, biophysics and physiology.
19
+
20
+ Args:
21
+ root (str): Root directory where the dataset should be saved.
22
+ name (str): The name of the dataset (:obj:`"ESOL"`, :obj:`"FreeSolv"`,
23
+ :obj:`"Lipo"`, :obj:`"PCBA"`, :obj:`"MUV"`, :obj:`"HIV"`,
24
+ :obj:`"BACE"`, :obj:`"BBBP"`, :obj:`"Tox21"`, :obj:`"ToxCast"`,
25
+ :obj:`"SIDER"`, :obj:`"ClinTox"`).
26
+ transform (callable, optional): A function/transform that takes in a
27
+ :obj:`k3_node.data.Data` object and returns a transformed version.
28
+ (default: :obj:`None`)
29
+ pre_transform (callable, optional): A function/transform that takes in a
30
+ :obj:`k3_node.data.Data` object and returns a transformed version.
31
+ (default: :obj:`None`)
32
+ pre_filter (callable, optional): A function that takes in a
33
+ :obj:`k3_node.data.Data` object and returns a boolean value,
34
+ indicating whether the data object should be included in the final
35
+ dataset. (default: :obj:`None`)
36
+ force_reload (bool, optional): Whether to re-process the dataset.
37
+ (default: :obj:`False`)
38
+ from_smiles (callable, optional): A custom function that takes a SMILES
39
+ string and outputs a :obj:`k3_node.data.Data` object.
40
+ (default: :obj:`None`)
41
+ """
42
+
43
+ url = "https://deepchemdata.s3-us-west-1.amazonaws.com/datasets/{}"
44
+
45
+ # Format: name: (display_name, url_name, csv_name, smiles_idx, y_idx)
46
+ names: Dict[str, Tuple[str, str, str, int, Union[int, slice]]] = {
47
+ "esol": ("ESOL", "delaney-processed.csv", "delaney-processed", -1, -2),
48
+ "freesolv": ("FreeSolv", "SAMPL.csv", "SAMPL", 1, 2),
49
+ "lipo": ("Lipophilicity", "Lipophilicity.csv", "Lipophilicity", 2, 1),
50
+ "pcba": ("PCBA", "pcba.csv.gz", "pcba", -1, slice(0, 128)),
51
+ "muv": ("MUV", "muv.csv.gz", "muv", -1, slice(0, 17)),
52
+ "hiv": ("HIV", "HIV.csv", "HIV", 0, -1),
53
+ "bace": ("BACE", "bace.csv", "bace", 0, 2),
54
+ "bbbp": ("BBBP", "BBBP.csv", "BBBP", -1, -2),
55
+ "tox21": ("Tox21", "tox21.csv.gz", "tox21", -1, slice(0, 12)),
56
+ "toxcast": ("ToxCast", "toxcast_data.csv.gz", "toxcast_data", 0, slice(1, 618)),
57
+ "sider": ("SIDER", "sider.csv.gz", "sider", 0, slice(1, 28)),
58
+ "clintox": ("ClinTox", "clintox.csv.gz", "clintox", 0, slice(1, 3)),
59
+ }
60
+
61
+ def __init__(
62
+ self,
63
+ root: str,
64
+ name: str,
65
+ transform: Optional[Callable] = None,
66
+ pre_transform: Optional[Callable] = None,
67
+ pre_filter: Optional[Callable] = None,
68
+ force_reload: bool = False,
69
+ from_smiles: Optional[Callable] = None,
70
+ ) -> None:
71
+ self.name = name.lower()
72
+ if self.name not in self.names:
73
+ raise ValueError(
74
+ f"Unknown dataset name '{name}'. Available names: {list(self.names.keys())}"
75
+ )
76
+ self.from_smiles = from_smiles or default_from_smiles
77
+ super().__init__(
78
+ root,
79
+ transform,
80
+ pre_transform,
81
+ pre_filter,
82
+ force_reload=force_reload,
83
+ )
84
+ self.load(self.processed_paths[0])
85
+
86
+ @property
87
+ def raw_dir(self) -> str:
88
+ return osp.join(self.root, self.name, "raw")
89
+
90
+ @property
91
+ def processed_dir(self) -> str:
92
+ return osp.join(self.root, self.name, "processed")
93
+
94
+ @property
95
+ def raw_file_names(self) -> str:
96
+ return f"{self.names[self.name][2]}.csv"
97
+
98
+ @property
99
+ def processed_file_names(self) -> str:
100
+ return "data.pt"
101
+
102
+ def download(self) -> None:
103
+ url = self.url.format(self.names[self.name][1])
104
+ path = download_url(url, self.raw_dir)
105
+ if self.names[self.name][1].endswith("gz"):
106
+ extract_gz(path, self.raw_dir)
107
+ os.unlink(path)
108
+
109
+ def process(self) -> None:
110
+ with open(self.raw_paths[0], "r", encoding="utf-8") as f:
111
+ dataset = f.read().split("\n")[1:-1]
112
+ dataset = [x for x in dataset if len(x) > 0]
113
+
114
+ data_list = []
115
+ for line in dataset:
116
+ line = re.sub(r'".*"', "", line) # Replace quoted substrings
117
+ values = line.split(",")
118
+
119
+ smiles = values[self.names[self.name][3]]
120
+ labels = values[self.names[self.name][4]]
121
+ labels = labels if isinstance(labels, list) else [labels]
122
+
123
+ ys = [float(y) if len(y) > 0 else float("nan") for y in labels]
124
+ y = ops.convert_to_tensor(np.array(ys, dtype=np.float32).reshape(1, -1), dtype="float32")
125
+
126
+ data = self.from_smiles(smiles)
127
+ data.y = y
128
+
129
+ if data.num_nodes == 0:
130
+ warnings.warn(
131
+ f"Skipping molecule '{smiles}' since it resulted in zero atoms",
132
+ stacklevel=2,
133
+ )
134
+ continue
135
+
136
+ if self.pre_filter is not None and not self.pre_filter(data):
137
+ continue
138
+
139
+ if self.pre_transform is not None:
140
+ data = self.pre_transform(data)
141
+
142
+ data_list.append(data)
143
+
144
+ self.save(data_list, self.processed_paths[0])
145
+
146
+ def __repr__(self) -> str:
147
+ return f"{self.names[self.name][0]}({len(self)})"
148
+
@@ -0,0 +1,7 @@
1
+ from .base import MotifGenerator
2
+ from .custom import CustomMotif
3
+ from .house import HouseMotif
4
+ from .cycle import CycleMotif
5
+
6
+ __all__ = ["MotifGenerator", "CustomMotif", "HouseMotif", "CycleMotif"]
7
+
@@ -0,0 +1,29 @@
1
+ from abc import ABC, abstractmethod
2
+ from typing import Any
3
+ from k3_node.data import Data
4
+
5
+
6
+ class MotifGenerator(ABC):
7
+ r"""An abstract base class for generating motifs."""
8
+
9
+ @abstractmethod
10
+ def __call__(self) -> Data:
11
+ raise NotImplementedError
12
+
13
+ @staticmethod
14
+ def resolve(query: Any, *args: Any, **kwargs: Any) -> "MotifGenerator":
15
+ if isinstance(query, MotifGenerator):
16
+ return query
17
+ if isinstance(query, str):
18
+ query = query.lower()
19
+ if query in ["house", "housemotif"]:
20
+ from k3_node.datasets.motif_generator.house import HouseMotif
21
+ return HouseMotif(*args, **kwargs)
22
+ elif query in ["cycle", "cyclemotif"]:
23
+ from k3_node.datasets.motif_generator.cycle import CycleMotif
24
+ return CycleMotif(*args, **kwargs)
25
+ raise ValueError(f"Could not resolve motif generator: {query}")
26
+
27
+ def __repr__(self) -> str:
28
+ return f"{self.__class__.__name__}()"
29
+
@@ -0,0 +1,17 @@
1
+ from typing import Any, Optional
2
+ from k3_node.data import Data
3
+ from k3_node.datasets.motif_generator.base import MotifGenerator
4
+
5
+
6
+ class CustomMotif(MotifGenerator):
7
+ r"""Generates a motif based on a custom structure coming from a Data object."""
8
+
9
+ def __init__(self, structure: Any):
10
+ super().__init__()
11
+ if not isinstance(structure, Data):
12
+ raise ValueError(f"Expected structure of type Data, got {type(structure)}")
13
+ self.structure = structure
14
+
15
+ def __call__(self) -> Data:
16
+ return self.structure
17
+
@@ -0,0 +1,25 @@
1
+ import numpy as np
2
+ from keras import ops
3
+
4
+ from k3_node.data import Data
5
+ from k3_node.datasets.motif_generator.custom import CustomMotif
6
+
7
+
8
+ class CycleMotif(CustomMotif):
9
+ r"""Generates the cycle motif from the "GNNExplainer" paper."""
10
+
11
+ def __init__(self, num_nodes: int):
12
+ self.num_nodes = num_nodes
13
+
14
+ row = np.repeat(np.arange(num_nodes), 2)
15
+ col1 = np.mod(np.arange(-1, num_nodes - 1), num_nodes)
16
+ col2 = np.mod(np.arange(1, num_nodes + 1), num_nodes)
17
+ col = np.sort(np.stack([col1, col2], axis=1), axis=-1).flatten()
18
+
19
+ edge_index = ops.convert_to_tensor(np.stack([row, col], axis=0), dtype="int64")
20
+ structure = Data(num_nodes=num_nodes, edge_index=edge_index)
21
+ super().__init__(structure)
22
+
23
+ def __repr__(self) -> str:
24
+ return f"{self.__class__.__name__}({self.num_nodes})"
25
+
@@ -0,0 +1,27 @@
1
+ import numpy as np
2
+ from keras import ops
3
+
4
+ from k3_node.data import Data
5
+ from k3_node.datasets.motif_generator.custom import CustomMotif
6
+
7
+
8
+ class HouseMotif(CustomMotif):
9
+ r"""Generates the house-structured motif from the "GNNExplainer" paper,
10
+ containing 5 nodes and 6 undirected edges.
11
+ """
12
+
13
+ def __init__(self):
14
+ edge_index = ops.convert_to_tensor(
15
+ np.array(
16
+ [
17
+ [0, 0, 0, 1, 1, 1, 2, 2, 3, 3, 4, 4],
18
+ [1, 3, 4, 4, 2, 0, 1, 3, 2, 0, 0, 1],
19
+ ],
20
+ dtype=np.int64,
21
+ ),
22
+ dtype="int64",
23
+ )
24
+ y = ops.convert_to_tensor(np.array([0, 0, 1, 1, 2], dtype=np.int64), dtype="int64")
25
+ structure = Data(num_nodes=5, edge_index=edge_index, y=y)
26
+ super().__init__(structure)
27
+
@@ -0,0 +1,55 @@
1
+ import os
2
+ import os.path as osp
3
+ from typing import Callable, Optional
4
+
5
+ import numpy as np
6
+
7
+ from k3_node.data import Data, InMemoryDataset
8
+ from k3_node.data.download import download_url
9
+ from k3_node.data.extract import extract_zip
10
+
11
+
12
+ class MovieLens100K(InMemoryDataset):
13
+ r"""The MovieLens 100K ratings (943 users, 1,682 movies, 100,000 ratings) as a user-movie graph
14
+ for recommendation, a small stand-in for PyG's ``AmazonBook``.
15
+
16
+ Every rating counts as an interaction. The graph is homogeneous: users are nodes
17
+ ``0 .. num_users - 1`` and movies ``num_users .. num_nodes - 1``. ``edge_index`` holds the
18
+ training interactions in both directions, ``edge_label_index`` the test interactions
19
+ (user, movie), from the official 80/20 split ``u1.base`` / ``u1.test``.
20
+
21
+ Args:
22
+ root (str): Root directory where the dataset should be saved.
23
+ transform (callable, optional): A function applied to the graph when it is accessed.
24
+ force_reload (bool, optional): Whether to re-process the dataset. (default: ``False``)
25
+ """
26
+
27
+ url = "https://files.grouplens.org/datasets/movielens/ml-100k.zip"
28
+ num_users, num_items = 943, 1682
29
+
30
+ def __init__(self, root: str, transform: Optional[Callable] = None, force_reload: bool = False):
31
+ super().__init__(root, transform, force_reload=force_reload)
32
+ self.load(self.processed_paths[0])
33
+
34
+ @property
35
+ def raw_file_names(self):
36
+ return [osp.join("ml-100k", "u1.base"), osp.join("ml-100k", "u1.test")]
37
+
38
+ @property
39
+ def processed_file_names(self) -> str:
40
+ return "data.pt"
41
+
42
+ def download(self):
43
+ path = download_url(self.url, self.raw_dir)
44
+ extract_zip(path, self.raw_dir)
45
+ os.unlink(path)
46
+
47
+ def _read(self, path):
48
+ ratings = np.loadtxt(path, dtype=np.int64)[:, :2] - 1 # user id, movie id (1-based)
49
+ return np.stack([ratings[:, 0], ratings[:, 1] + self.num_users])
50
+
51
+ def process(self):
52
+ train, test = self._read(self.raw_paths[0]), self._read(self.raw_paths[1])
53
+ data = Data(edge_index=np.concatenate([train, train[::-1]], axis=1), edge_label_index=test,
54
+ num_nodes=self.num_users + self.num_items)
55
+ self.save([data], self.processed_paths[0])