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,256 @@
1
+ import random
2
+ from collections import defaultdict
3
+ from itertools import product
4
+ from typing import Callable, Dict, List, Optional, Tuple, Union
5
+ import numpy as np
6
+ from keras import ops
7
+
8
+ from k3_node.data import Data, HeteroData, InMemoryDataset
9
+ from k3_node.layers.conv.utils import remove_self_loops
10
+ from k3_node.transforms.utils import to_undirected
11
+ from k3_node.utils.graph import coalesce
12
+
13
+
14
+ def get_num_nodes(avg_num_nodes: int, avg_degree: float) -> int:
15
+ min_num_nodes = max(3 * avg_num_nodes // 4, int(avg_degree))
16
+ max_num_nodes = 5 * avg_num_nodes // 4
17
+ return random.randint(min_num_nodes, max_num_nodes)
18
+
19
+
20
+ def get_num_channels(num_channels: int) -> int:
21
+ min_num_channels = 3 * num_channels // 4
22
+ max_num_channels = 5 * num_channels // 4
23
+ return random.randint(min_num_channels, max_num_channels)
24
+
25
+
26
+ def get_edge_index(
27
+ num_src_nodes: int,
28
+ num_dst_nodes: int,
29
+ avg_degree: float,
30
+ is_undirected: bool = False,
31
+ remove_loops: bool = False,
32
+ ):
33
+ num_edges = int(num_src_nodes * avg_degree)
34
+ row = np.random.randint(0, num_src_nodes, size=(num_edges,), dtype=np.int64)
35
+ col = np.random.randint(0, num_dst_nodes, size=(num_edges,), dtype=np.int64)
36
+ edge_index = np.stack([row, col], axis=0)
37
+
38
+ if remove_loops:
39
+ edge_index, _ = remove_self_loops(edge_index)
40
+
41
+ num_nodes = max(num_src_nodes, num_dst_nodes)
42
+ if is_undirected:
43
+ edge_index = to_undirected(edge_index, num_nodes=num_nodes)
44
+ else:
45
+ edge_index, _ = coalesce(edge_index, num_nodes=num_nodes)
46
+
47
+ return ops.convert_to_tensor(edge_index, dtype="int64")
48
+
49
+
50
+ class FakeDataset(InMemoryDataset):
51
+ r"""A fake dataset that returns randomly generated `k3_node.data.Data` objects."""
52
+
53
+ def __init__(
54
+ self,
55
+ num_graphs: int = 1,
56
+ avg_num_nodes: int = 1000,
57
+ avg_degree: float = 10.0,
58
+ num_channels: int = 64,
59
+ edge_dim: int = 0,
60
+ num_classes: int = 10,
61
+ task: str = "auto",
62
+ is_undirected: bool = True,
63
+ transform: Optional[Callable] = None,
64
+ pre_transform: Optional[Callable] = None,
65
+ **kwargs: Union[int, Tuple[int, ...]],
66
+ ):
67
+ super().__init__(None, transform)
68
+
69
+ if task == "auto":
70
+ task = "graph" if num_graphs > 1 else "node"
71
+ assert task in ["node", "graph"]
72
+
73
+ self.num_graphs_val = num_graphs
74
+ self.avg_num_nodes = max(avg_num_nodes, int(avg_degree))
75
+ self.avg_degree = max(avg_degree, 1)
76
+ self.num_channels = num_channels
77
+ self.edge_dim = edge_dim
78
+ self._num_classes = num_classes
79
+ self.task = task
80
+ self.is_undirected = is_undirected
81
+ self.kwargs = kwargs
82
+
83
+ data_list = [self.generate_data() for _ in range(max(num_graphs, 1))]
84
+ self.data, self.slices = self.collate(data_list)
85
+
86
+ def __repr__(self) -> str:
87
+ return f"FakeDataset({self.num_graphs_val})" if self.num_graphs_val > 1 else "FakeDataset()"
88
+
89
+ def generate_data(self) -> Data:
90
+ num_nodes = get_num_nodes(self.avg_num_nodes, self.avg_degree)
91
+ data = Data()
92
+
93
+ if self._num_classes > 0 and self.task == "node":
94
+ data.y = ops.convert_to_tensor(
95
+ np.random.randint(0, self._num_classes, size=(num_nodes,), dtype=np.int64),
96
+ dtype="int64",
97
+ )
98
+ elif self._num_classes > 0 and self.task == "graph":
99
+ data.y = ops.convert_to_tensor(
100
+ np.array([random.randint(0, self._num_classes - 1)], dtype=np.int64),
101
+ dtype="int64",
102
+ )
103
+
104
+ data.edge_index = get_edge_index(
105
+ num_nodes, num_nodes, self.avg_degree, self.is_undirected, remove_loops=True
106
+ )
107
+
108
+ if self.num_channels > 0:
109
+ x = np.random.randn(num_nodes, self.num_channels).astype(np.float32)
110
+ if self._num_classes > 0 and self.task == "node":
111
+ y_np = ops.convert_to_numpy(data.y)
112
+ x = x + y_np[:, None]
113
+ elif self._num_classes > 0 and self.task == "graph":
114
+ y_np = ops.convert_to_numpy(data.y)
115
+ x = x + y_np
116
+ data.x = ops.convert_to_tensor(x, dtype="float32")
117
+ else:
118
+ data.num_nodes = num_nodes
119
+
120
+ num_edges = int(ops.shape(data.edge_index)[1])
121
+ if self.edge_dim > 1:
122
+ data.edge_attr = ops.convert_to_tensor(
123
+ np.random.rand(num_edges, self.edge_dim).astype(np.float32),
124
+ dtype="float32",
125
+ )
126
+ elif self.edge_dim == 1:
127
+ data.edge_weight = ops.convert_to_tensor(
128
+ np.random.rand(num_edges).astype(np.float32),
129
+ dtype="float32",
130
+ )
131
+
132
+ for feature_name, feature_shape in self.kwargs.items():
133
+ shape = (feature_shape,) if isinstance(feature_shape, int) else feature_shape
134
+ setattr(
135
+ data,
136
+ feature_name,
137
+ ops.convert_to_tensor(np.random.randn(*shape).astype(np.float32), dtype="float32"),
138
+ )
139
+
140
+ return data
141
+
142
+
143
+ class FakeHeteroDataset(InMemoryDataset):
144
+ r"""A fake dataset that returns randomly generated `k3_node.data.HeteroData` objects."""
145
+
146
+ def __init__(
147
+ self,
148
+ num_graphs: int = 1,
149
+ num_node_types: int = 3,
150
+ num_edge_types: int = 6,
151
+ avg_num_nodes: int = 1000,
152
+ avg_degree: float = 10.0,
153
+ avg_num_channels: int = 64,
154
+ edge_dim: int = 0,
155
+ num_classes: int = 10,
156
+ task: str = "auto",
157
+ transform: Optional[Callable] = None,
158
+ pre_transform: Optional[Callable] = None,
159
+ **kwargs: Union[int, Tuple[int, ...]],
160
+ ):
161
+ super().__init__(None, transform)
162
+
163
+ if task == "auto":
164
+ task = "graph" if num_graphs > 1 else "node"
165
+ assert task in ["node", "graph"]
166
+
167
+ self.num_graphs_val = num_graphs
168
+ self.node_types = [f"v{i}" for i in range(max(num_node_types, 1))]
169
+
170
+ edge_types: List[Tuple[str, str]] = []
171
+ edge_type_product = list(product(self.node_types, self.node_types))
172
+ while len(edge_types) < max(num_edge_types, 1):
173
+ edge_types.extend(edge_type_product)
174
+ random.shuffle(edge_types)
175
+
176
+ self.edge_types: List[Tuple[str, str, str]] = []
177
+ count: Dict[Tuple[str, str], int] = defaultdict(int)
178
+ for edge_type in edge_types[: max(num_edge_types, 1)]:
179
+ rel = f"e{count[edge_type]}"
180
+ count[edge_type] += 1
181
+ self.edge_types.append((edge_type[0], rel, edge_type[1]))
182
+
183
+ self.avg_num_nodes = max(avg_num_nodes, int(avg_degree))
184
+ self.avg_degree = max(avg_degree, 1)
185
+ self.avg_num_channels = avg_num_channels
186
+ self.edge_dim = edge_dim
187
+ self._num_classes = num_classes
188
+ self.task = task
189
+ self.kwargs = kwargs
190
+
191
+ data_list = [self.generate_data() for _ in range(max(num_graphs, 1))]
192
+ self.data, self.slices = self.collate(data_list)
193
+
194
+ def __repr__(self) -> str:
195
+ return f"FakeHeteroDataset({self.num_graphs_val})" if self.num_graphs_val > 1 else "FakeHeteroDataset()"
196
+
197
+ def generate_data(self) -> HeteroData:
198
+ data = HeteroData()
199
+
200
+ for node_type in self.node_types:
201
+ num_nodes = get_num_nodes(self.avg_num_nodes, self.avg_degree)
202
+ num_channels = get_num_channels(self.avg_num_channels)
203
+ store = data[node_type]
204
+
205
+ if self.avg_num_channels > 0:
206
+ store.x = ops.convert_to_tensor(
207
+ np.random.randn(num_nodes, num_channels).astype(np.float32),
208
+ dtype="float32",
209
+ )
210
+ else:
211
+ store.num_nodes = num_nodes
212
+
213
+ if self._num_classes > 0 and self.task == "node":
214
+ store.y = ops.convert_to_tensor(
215
+ np.random.randint(0, self._num_classes, size=(num_nodes,), dtype=np.int64),
216
+ dtype="int64",
217
+ )
218
+
219
+ for edge_type in self.edge_types:
220
+ src, rel, dst = edge_type
221
+ store = data[edge_type]
222
+ store.edge_index = get_edge_index(
223
+ data[src].num_nodes,
224
+ data[dst].num_nodes,
225
+ self.avg_degree,
226
+ is_undirected=False,
227
+ remove_loops=False,
228
+ )
229
+
230
+ num_edges = int(ops.shape(store.edge_index)[1])
231
+ if self.edge_dim > 1:
232
+ store.edge_attr = ops.convert_to_tensor(
233
+ np.random.rand(num_edges, self.edge_dim).astype(np.float32),
234
+ dtype="float32",
235
+ )
236
+ elif self.edge_dim == 1:
237
+ store.edge_weight = ops.convert_to_tensor(
238
+ np.random.rand(num_edges).astype(np.float32),
239
+ dtype="float32",
240
+ )
241
+
242
+ if self._num_classes > 0 and self.task == "graph":
243
+ data.y = ops.convert_to_tensor(
244
+ np.array([random.randint(0, self._num_classes - 1)], dtype=np.int64),
245
+ dtype="int64",
246
+ )
247
+
248
+ for feature_name, feature_shape in self.kwargs.items():
249
+ shape = (feature_shape,) if isinstance(feature_shape, int) else feature_shape
250
+ setattr(
251
+ data,
252
+ feature_name,
253
+ ops.convert_to_tensor(np.random.randn(*shape).astype(np.float32), dtype="float32"),
254
+ )
255
+
256
+ return data
@@ -0,0 +1,90 @@
1
+ from typing import Callable, List, 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 FB15k_237(InMemoryDataset):
10
+ r"""The FB15K237 dataset containing 14,541 entities, 237 relations and 310,116 fact triples.
11
+
12
+ Args:
13
+ root (str): Root directory where the dataset should be saved.
14
+ split (str, optional): "train", "val", or "test". (default: "train")
15
+ transform (callable, optional): Transform function.
16
+ pre_transform (callable, optional): Pre-transform function.
17
+ force_reload (bool, optional): Whether to re-process the dataset.
18
+ """
19
+
20
+ url = "https://raw.githubusercontent.com/villmow/datasets_knowledge_embedding/master/FB15k-237"
21
+
22
+ def __init__(
23
+ self,
24
+ root: str,
25
+ split: str = "train",
26
+ transform: Optional[Callable] = None,
27
+ pre_transform: Optional[Callable] = None,
28
+ force_reload: bool = False,
29
+ ):
30
+ if split not in {"train", "val", "test"}:
31
+ raise ValueError(f"Invalid split argument (got {split})")
32
+ self.split = split
33
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
34
+ idx = ["train", "val", "test"].index(split)
35
+ self.load(self.processed_paths[idx])
36
+
37
+ @property
38
+ def raw_file_names(self) -> List[str]:
39
+ return ["train.txt", "valid.txt", "test.txt"]
40
+
41
+ @property
42
+ def processed_file_names(self) -> List[str]:
43
+ return ["train_data.pt", "val_data.pt", "test_data.pt"]
44
+
45
+ def download(self):
46
+ for filename in self.raw_file_names:
47
+ fs.cp(f"{self.url}/{filename}", self.raw_dir)
48
+
49
+ def process(self):
50
+ # Map entities and relations to integer IDs
51
+ entities, relations = {}, {}
52
+ for path in self.raw_paths:
53
+ with open(path) as f:
54
+ lines = f.read().split("\n")[:-1]
55
+ for line in lines:
56
+ parts = line.split()
57
+ if len(parts) >= 3:
58
+ s, r, d = parts[0], parts[1], parts[2]
59
+ if s not in entities:
60
+ entities[s] = len(entities)
61
+ if d not in entities:
62
+ entities[d] = len(entities)
63
+ if r not in relations:
64
+ relations[r] = len(relations)
65
+
66
+ for in_path, out_path in zip(self.raw_paths, self.processed_paths):
67
+ srcs, dsts, rels = [], [], []
68
+ with open(in_path) as f:
69
+ lines = f.read().split("\n")[:-1]
70
+ for line in lines:
71
+ parts = line.split()
72
+ if len(parts) >= 3:
73
+ srcs.append(entities[parts[0]])
74
+ rels.append(relations[parts[1]])
75
+ dsts.append(entities[parts[2]])
76
+
77
+ edge_index = np.array([srcs, dsts], dtype=np.int64)
78
+ edge_type = np.array(rels, dtype=np.int64)
79
+
80
+ data = Data(
81
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
82
+ edge_type=ops.convert_to_tensor(edge_type, dtype="int64"),
83
+ num_nodes=len(entities),
84
+ )
85
+
86
+ if self.pre_transform is not None:
87
+ data = self.pre_transform(data)
88
+
89
+ self.save([data], out_path)
90
+
@@ -0,0 +1,69 @@
1
+ import glob
2
+ import os
3
+ import os.path as osp
4
+ from typing import Callable, List, Optional
5
+
6
+ import numpy as np
7
+
8
+ from k3_node.data import InMemoryDataset
9
+ from k3_node.data.download import download_url
10
+ from k3_node.data.extract import extract_zip
11
+ from k3_node.io.off import read_off
12
+
13
+
14
+ class GeometricShapes(InMemoryDataset):
15
+ r"""Synthetic meshes of 40 geometric shapes such as cubes, spheres or pyramids (one training and
16
+ one test mesh per shape), as in PyG.
17
+
18
+ The graphs hold mesh faces (``face``) and vertex positions (``pos``) but no edges. Use
19
+ :class:`~k3_node.transforms.FaceToEdge` to turn a mesh into a graph, or
20
+ :class:`~k3_node.transforms.SamplePoints` to sample a point cloud from its surface.
21
+
22
+ Args:
23
+ root (str): Root directory where the dataset should be saved.
24
+ train (bool, optional): Loads the training meshes if ``True``, else the test meshes.
25
+ transform (callable, optional): A function applied to each graph when it is accessed.
26
+ pre_transform (callable, optional): A function applied to each graph before saving.
27
+ pre_filter (callable, optional): A function deciding which graphs to keep.
28
+ force_reload (bool, optional): Whether to re-process the dataset. (default: ``False``)
29
+ """
30
+
31
+ url = 'https://github.com/Yannick-S/geometric_shapes/raw/master/raw.zip'
32
+
33
+ def __init__(self, root: str, train: bool = True, transform: Optional[Callable] = None,
34
+ pre_transform: Optional[Callable] = None, pre_filter: Optional[Callable] = None,
35
+ force_reload: bool = False):
36
+ super().__init__(root, transform, pre_transform, pre_filter, force_reload=force_reload)
37
+ self.load(self.processed_paths[0] if train else self.processed_paths[1])
38
+
39
+ @property
40
+ def raw_file_names(self) -> str:
41
+ return '2d_circle'
42
+
43
+ @property
44
+ def processed_file_names(self) -> List[str]:
45
+ return ['training.pt', 'test.pt']
46
+
47
+ def download(self):
48
+ path = download_url(self.url, self.root)
49
+ extract_zip(path, self.root)
50
+ os.unlink(path)
51
+
52
+ def process(self):
53
+ self.save(self._process_set('train'), self.processed_paths[0])
54
+ self.save(self._process_set('test'), self.processed_paths[1])
55
+
56
+ def _process_set(self, split: str):
57
+ categories = sorted(x.split(os.sep)[-2] for x in glob.glob(osp.join(self.raw_dir, '*', '')))
58
+ data_list = []
59
+ for target, category in enumerate(categories):
60
+ for path in sorted(glob.glob(osp.join(self.raw_dir, category, split, '*.off'))):
61
+ data = read_off(path)
62
+ data.pos = data.pos - data.pos.mean(axis=0, keepdims=True)
63
+ data.y = np.array([target], dtype=np.int64)
64
+ data_list.append(data)
65
+ if self.pre_filter is not None:
66
+ data_list = [d for d in data_list if self.pre_filter(d)]
67
+ if self.pre_transform is not None:
68
+ data_list = [self.pre_transform(d) for d in data_list]
69
+ return data_list
@@ -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 GitHub(InMemoryDataset):
10
+ r"""The GitHub Web and ML Developers dataset."""
11
+
12
+ url = "https://graphmining.ai/datasets/ptg/github.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 "github.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,6 @@
1
+ from .base import GraphGenerator
2
+ from .ba_graph import BAGraph
3
+ from .er_graph import ERGraph
4
+
5
+ __all__ = ["GraphGenerator", "BAGraph", "ERGraph"]
6
+
@@ -0,0 +1,20 @@
1
+ from k3_node.data import Data
2
+ from k3_node.datasets.graph_generator.base import GraphGenerator
3
+ from k3_node.utils.random import barabasi_albert_graph
4
+
5
+
6
+ class BAGraph(GraphGenerator):
7
+ r"""Generates random Barabasi-Albert (BA) graphs."""
8
+
9
+ def __init__(self, num_nodes: int, num_edges: int):
10
+ super().__init__()
11
+ self.num_nodes = num_nodes
12
+ self.num_edges = num_edges
13
+
14
+ def __call__(self) -> Data:
15
+ edge_index = barabasi_albert_graph(self.num_nodes, self.num_edges)
16
+ return Data(num_nodes=self.num_nodes, edge_index=edge_index)
17
+
18
+ def __repr__(self) -> str:
19
+ return f"{self.__class__.__name__}(num_nodes={self.num_nodes}, num_edges={self.num_edges})"
20
+
@@ -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 GraphGenerator(ABC):
7
+ r"""An abstract base class for generating synthetic graphs."""
8
+
9
+ @abstractmethod
10
+ def __call__(self) -> Data:
11
+ raise NotImplementedError
12
+
13
+ @staticmethod
14
+ def resolve(query: Any, *args: Any, **kwargs: Any) -> "GraphGenerator":
15
+ if isinstance(query, GraphGenerator):
16
+ return query
17
+ if isinstance(query, str):
18
+ query = query.lower()
19
+ if query in ["ba", "bagraph", "barabasi_albert"]:
20
+ from k3_node.datasets.graph_generator.ba_graph import BAGraph
21
+ return BAGraph(*args, **kwargs)
22
+ elif query in ["er", "ergraph", "erdos_renyi"]:
23
+ from k3_node.datasets.graph_generator.er_graph import ERGraph
24
+ return ERGraph(*args, **kwargs)
25
+ raise ValueError(f"Could not resolve graph generator: {query}")
26
+
27
+ def __repr__(self) -> str:
28
+ return f"{self.__class__.__name__}()"
29
+
@@ -0,0 +1,21 @@
1
+ from k3_node.data import Data
2
+ from k3_node.datasets.graph_generator.base import GraphGenerator
3
+ from k3_node.utils.random import erdos_renyi_graph
4
+
5
+
6
+ class ERGraph(GraphGenerator):
7
+ r"""Generates random Erdos-Renyi (ER) graphs."""
8
+
9
+ def __init__(self, num_nodes: int, edge_prob: float, directed: bool = False):
10
+ super().__init__()
11
+ self.num_nodes = num_nodes
12
+ self.edge_prob = edge_prob
13
+ self.directed = directed
14
+
15
+ def __call__(self) -> Data:
16
+ edge_index = erdos_renyi_graph(self.num_nodes, self.edge_prob, directed=self.directed)
17
+ return Data(num_nodes=self.num_nodes, edge_index=edge_index)
18
+
19
+ def __repr__(self) -> str:
20
+ return f"{self.__class__.__name__}(num_nodes={self.num_nodes}, edge_prob={self.edge_prob})"
21
+
@@ -0,0 +1,58 @@
1
+ from typing import Callable, List, Optional
2
+
3
+ import numpy as np
4
+
5
+ from k3_node.data import Data, InMemoryDataset
6
+ from k3_node.data.download import download_url
7
+
8
+
9
+ class ICEWS18(InMemoryDataset):
10
+ r"""The ICEWS18 temporal knowledge graph (Integrated Crisis Early Warning System, events from
11
+ 1/1/2018 to 10/31/2018 at a daily resolution), used by RE-Net: every graph is one event with
12
+ subject ``sub``, relation ``rel``, object ``obj`` and day ``t``.
13
+
14
+ Args:
15
+ root (str): Root directory where the dataset should be saved.
16
+ split (str, optional): ``"train"``, ``"val"`` or ``"test"``. (default: ``"train"``)
17
+ transform (callable, optional): A function applied to each event when it is accessed.
18
+ pre_transform (callable, optional): A function applied to the events in time order before
19
+ saving, e.g. :meth:`~k3_node.models.RENet.pre_transform`.
20
+ force_reload (bool, optional): Whether to re-process the dataset. (default: ``False``)
21
+ """
22
+
23
+ url = 'https://github.com/INK-USC/RE-Net/raw/master/data/ICEWS18'
24
+ splits = [0, 373018, 419013, 468558]
25
+ num_nodes, num_rels = 23033, 256
26
+
27
+ def __init__(self, root: str, split: str = 'train', transform: Optional[Callable] = None,
28
+ pre_transform: Optional[Callable] = None, force_reload: bool = False):
29
+ assert split in ['train', 'val', 'test']
30
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
31
+ self.load(self.processed_paths[['train', 'val', 'test'].index(split)])
32
+
33
+ @property
34
+ def raw_file_names(self) -> List[str]:
35
+ return [f'{name}.txt' for name in ['train', 'valid', 'test']]
36
+
37
+ @property
38
+ def processed_file_names(self) -> List[str]:
39
+ return ['train.pt', 'val.pt', 'test.pt']
40
+
41
+ def download(self):
42
+ for filename in self.raw_file_names:
43
+ download_url(f'{self.url}/{filename}', self.raw_dir)
44
+
45
+ def process(self):
46
+ events = np.concatenate([np.loadtxt(path, delimiter='\t', usecols=range(4), dtype=np.int64)
47
+ for path in self.raw_paths])
48
+ events[:, 3] //= 24 # hours -> days
49
+ events -= events.min(axis=0, keepdims=True)
50
+ data_list = []
51
+ for sub, rel, obj, t in events.tolist():
52
+ data = Data(sub=sub, rel=rel, obj=obj, t=t)
53
+ if self.pre_transform is not None:
54
+ data = self.pre_transform(data)
55
+ data_list.append(data)
56
+ s = self.splits
57
+ for i in range(3):
58
+ self.save(data_list[s[i]:s[i + 1]], self.processed_paths[i])