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,322 @@
1
+ import os
2
+ import os.path as osp
3
+ import pickle
4
+ import tempfile
5
+ import numpy as np
6
+ import pytest
7
+ import scipy.sparse as sp
8
+ from keras import ops
9
+
10
+ import k3_node.datasets as datasets
11
+ import k3_node.io as io
12
+
13
+
14
+ def to_np(x):
15
+ if x is None:
16
+ return None
17
+ return ops.convert_to_numpy(x)
18
+
19
+
20
+ def test_karate_club():
21
+ dataset = datasets.KarateClub()
22
+ assert len(dataset) == 1
23
+ assert repr(dataset) == "KarateClub()"
24
+ assert dataset.num_classes == 4
25
+ assert dataset.num_features == 34
26
+ assert dataset.num_node_features == 34
27
+
28
+ data = dataset[0]
29
+ assert data.num_nodes == 34
30
+ assert data.num_edges == 156
31
+ assert data.x.shape == (34, 34)
32
+ assert data.edge_index.shape == (2, 156)
33
+ assert data.y.shape == (34,)
34
+ assert data.train_mask.shape == (34,)
35
+
36
+ # Exactly 4 train nodes (1 per class)
37
+ assert int(np.sum(to_np(data.train_mask))) == 4
38
+ assert data.is_undirected() is True
39
+ assert data.has_self_loops() is False
40
+
41
+
42
+ def test_fake_dataset_node():
43
+ dataset = datasets.FakeDataset(
44
+ num_graphs=1,
45
+ avg_num_nodes=50,
46
+ avg_degree=5.0,
47
+ num_channels=16,
48
+ edge_dim=4,
49
+ num_classes=5,
50
+ task="node",
51
+ )
52
+ assert len(dataset) == 1
53
+ assert repr(dataset) == "FakeDataset()"
54
+ data = dataset[0]
55
+ assert data.num_nodes > 0
56
+ assert data.num_edges > 0
57
+ assert data.x.shape[-1] == 16
58
+ assert data.edge_attr.shape[-1] == 4
59
+ assert data.y.shape == (data.num_nodes,)
60
+
61
+
62
+ def test_fake_dataset_graph():
63
+ dataset = datasets.FakeDataset(
64
+ num_graphs=3,
65
+ avg_num_nodes=30,
66
+ avg_degree=4.0,
67
+ num_channels=8,
68
+ num_classes=2,
69
+ task="graph",
70
+ )
71
+ assert len(dataset) == 3
72
+ assert repr(dataset) == "FakeDataset(3)"
73
+ for i in range(3):
74
+ data = dataset[i]
75
+ assert data.num_nodes > 0
76
+ assert data.x.shape[-1] == 8
77
+ assert data.y.shape == (1,)
78
+
79
+
80
+ def test_fake_hetero_dataset():
81
+ dataset = datasets.FakeHeteroDataset(
82
+ num_graphs=1,
83
+ num_node_types=2,
84
+ num_edge_types=2,
85
+ avg_num_nodes=30,
86
+ avg_degree=3.0,
87
+ avg_num_channels=8,
88
+ num_classes=3,
89
+ task="node",
90
+ )
91
+ assert len(dataset) == 1
92
+ data = dataset[0]
93
+ assert len(data.node_types) == 2
94
+ assert len(data.edge_types) == 2
95
+
96
+ for nt in data.node_types:
97
+ assert data[nt].num_nodes > 0
98
+ assert data[nt].x.shape[-1] > 0
99
+ assert data[nt].y.shape == (data[nt].num_nodes,)
100
+
101
+ for et in data.edge_types:
102
+ assert data[et].edge_index.shape[0] == 2
103
+
104
+
105
+ def test_ba_shapes():
106
+ dataset = datasets.BAShapes(connection_distribution="random")
107
+ assert len(dataset) == 1
108
+ data = dataset[0]
109
+ assert data.num_nodes == 700
110
+ assert data.x.shape == (700, 10)
111
+ assert data.expl_mask.shape == (700,)
112
+ assert int(to_np(data.y).max()) == 3
113
+
114
+
115
+ def test_sbm_dataset():
116
+ with tempfile.TemporaryDirectory() as tmp_dir:
117
+ dataset = datasets.StochasticBlockModelDataset(
118
+ root=tmp_dir,
119
+ block_sizes=[10, 10],
120
+ edge_probs=[[0.5, 0.05], [0.05, 0.5]],
121
+ num_graphs=2,
122
+ num_channels=4,
123
+ )
124
+ assert len(dataset) == 2
125
+ d0 = dataset[0]
126
+ assert d0.num_nodes == 20
127
+ assert d0.x.shape == (20, 4)
128
+ assert d0.y.shape == (20,)
129
+
130
+
131
+ def test_explainer_dataset():
132
+ dataset = datasets.ExplainerDataset(
133
+ graph_generator="ba",
134
+ motif_generator="house",
135
+ num_motifs=3,
136
+ num_graphs=2,
137
+ graph_generator_kwargs={"num_nodes": 30, "num_edges": 2},
138
+ )
139
+ assert len(dataset) == 2
140
+ for i in range(2):
141
+ d = dataset[i]
142
+ assert d.edge_index.shape[0] == 2
143
+ assert d.node_mask.shape[0] == d.num_nodes
144
+ assert d.edge_mask.shape[0] == d.edge_index.shape[1]
145
+
146
+
147
+ def test_read_npz():
148
+ with tempfile.TemporaryDirectory() as tmp_dir:
149
+ path = osp.join(tmp_dir, "test.npz")
150
+ x_dense = (np.random.rand(8, 4) > 0.5).astype(np.float32)
151
+ x_csr = sp.csr_matrix(x_dense)
152
+
153
+ adj_dense = np.zeros((8, 8), dtype=np.float32)
154
+ adj_dense[0, 1] = adj_dense[1, 0] = 1.0
155
+ adj_csr = sp.csr_matrix(adj_dense)
156
+ labels = np.array([0, 1, 0, 1, 0, 1, 0, 1], dtype=np.int64)
157
+
158
+ np.savez(
159
+ path,
160
+ attr_data=x_csr.data,
161
+ attr_indices=x_csr.indices,
162
+ attr_indptr=x_csr.indptr,
163
+ attr_shape=x_csr.shape,
164
+ adj_data=adj_csr.data,
165
+ adj_indices=adj_csr.indices,
166
+ adj_indptr=adj_csr.indptr,
167
+ adj_shape=adj_csr.shape,
168
+ labels=labels,
169
+ )
170
+
171
+ data = io.read_npz(path, to_undirected=True)
172
+ assert data.num_nodes == 8
173
+ assert data.x.shape == (8, 4)
174
+ assert data.edge_index.shape[0] == 2
175
+ assert data.y.shape == (8,)
176
+
177
+
178
+ def test_read_planetoid():
179
+ with tempfile.TemporaryDirectory() as tmp_dir:
180
+ prefix = "cora"
181
+ num_train = 10
182
+ num_val = 500
183
+ num_test = 20
184
+ num_nodes = num_train + num_val + num_test
185
+
186
+ x = sp.csr_matrix(np.random.randn(num_train, 6).astype(np.float32))
187
+ y = np.eye(2)[np.random.randint(0, 2, num_train)]
188
+ tx = sp.csr_matrix(np.random.randn(num_test, 6).astype(np.float32))
189
+ ty = np.eye(2)[np.random.randint(0, 2, num_test)]
190
+ allx = sp.csr_matrix(np.random.randn(num_train + num_val, 6).astype(np.float32))
191
+ ally = np.eye(2)[np.random.randint(0, 2, num_train + num_val)]
192
+ graph = {i: [(i + 1) % num_nodes] for i in range(num_nodes)}
193
+ test_index = np.arange(num_train + num_val, num_nodes, dtype=np.int64)
194
+
195
+ for name, obj in [
196
+ ("x", x),
197
+ ("tx", tx),
198
+ ("allx", allx),
199
+ ("y", y),
200
+ ("ty", ty),
201
+ ("ally", ally),
202
+ ("graph", graph),
203
+ ]:
204
+ with open(osp.join(tmp_dir, f"ind.{prefix}.{name}"), "wb") as f:
205
+ pickle.dump(obj, f, protocol=2)
206
+
207
+ np.savetxt(osp.join(tmp_dir, f"ind.{prefix}.test.index"), test_index, fmt="%d")
208
+
209
+ data = io.read_planetoid_data(tmp_dir, prefix)
210
+ assert data.num_nodes == num_nodes
211
+ assert data.x.shape == (num_nodes, 6)
212
+ assert data.y.shape == (num_nodes,)
213
+ assert data.train_mask.shape == (num_nodes,)
214
+ assert data.test_mask.shape == (num_nodes,)
215
+
216
+
217
+ def test_read_tu():
218
+ with tempfile.TemporaryDirectory() as tmp_dir:
219
+ prefix = "MUTAG"
220
+ np.savetxt(osp.join(tmp_dir, f"{prefix}_graph_indicator.txt"), [1, 1, 2, 2], fmt="%d")
221
+ np.savetxt(
222
+ osp.join(tmp_dir, f"{prefix}_A.txt"),
223
+ [[1, 2], [2, 1], [3, 4], [4, 3]],
224
+ fmt="%d",
225
+ delimiter=", ",
226
+ )
227
+ np.savetxt(osp.join(tmp_dir, f"{prefix}_node_labels.txt"), [0, 1, 0, 1], fmt="%d")
228
+ np.savetxt(osp.join(tmp_dir, f"{prefix}_graph_labels.txt"), [1, 0], fmt="%d")
229
+
230
+ data, slices, sizes = io.read_tu_data(tmp_dir, prefix)
231
+ assert data.num_nodes == 4
232
+ assert "edge_index" in slices
233
+ assert "x" in slices
234
+ assert "y" in slices
235
+ assert sizes["num_node_labels"] == 2
236
+
237
+
238
+ def test_ppi():
239
+ import json
240
+ with tempfile.TemporaryDirectory() as tmp_dir:
241
+ raw_dir = osp.join(tmp_dir, "raw")
242
+ os.makedirs(raw_dir, exist_ok=True)
243
+
244
+ for split in ["train", "valid", "test"]:
245
+ graph = {
246
+ "directed": True,
247
+ "multigraph": False,
248
+ "graph": {},
249
+ "nodes": [{"id": 0}, {"id": 1}, {"id": 2}, {"id": 3}],
250
+ "links": [
251
+ {"source": 0, "target": 1},
252
+ {"source": 1, "target": 0},
253
+ {"source": 2, "target": 3},
254
+ {"source": 3, "target": 2},
255
+ ],
256
+ }
257
+ with open(osp.join(raw_dir, f"{split}_graph.json"), "w") as f:
258
+ json.dump(graph, f)
259
+
260
+ np.save(osp.join(raw_dir, f"{split}_feats.npy"), np.ones((4, 50), dtype=np.float32))
261
+ np.save(osp.join(raw_dir, f"{split}_labels.npy"), np.ones((4, 121), dtype=np.float32))
262
+ np.save(osp.join(raw_dir, f"{split}_graph_id.npy"), np.array([1, 1, 2, 2], dtype=np.int64))
263
+
264
+ dataset = datasets.PPI(root=tmp_dir, split="train")
265
+ assert len(dataset) == 2
266
+ assert dataset.num_features == 50
267
+ assert dataset.num_classes == 121
268
+
269
+ data0 = dataset[0]
270
+ assert data0.num_nodes == 2
271
+ assert data0.x.shape == (2, 50)
272
+ assert data0.y.shape == (2, 121)
273
+ assert data0.edge_index.shape == (2, 2)
274
+
275
+ val_dataset = datasets.PPI(root=tmp_dir, split="val")
276
+ assert len(val_dataset) == 2
277
+
278
+ with pytest.raises(AssertionError):
279
+ datasets.PPI(root=tmp_dir, split="unknown")
280
+
281
+
282
+ def test_reddit():
283
+ with tempfile.TemporaryDirectory() as tmp_dir:
284
+ raw_dir = osp.join(tmp_dir, "raw")
285
+ os.makedirs(raw_dir, exist_ok=True)
286
+
287
+ num_nodes = 10
288
+ num_features = 602
289
+ features = np.random.randn(num_nodes, num_features).astype(np.float32)
290
+ labels = np.random.randint(0, 41, size=num_nodes, dtype=np.int64)
291
+ node_types = np.array([1, 1, 1, 1, 2, 2, 2, 3, 3, 3], dtype=np.int64)
292
+
293
+ np.savez(
294
+ osp.join(raw_dir, "reddit_data.npz"),
295
+ feature=features,
296
+ label=labels,
297
+ node_types=node_types,
298
+ )
299
+
300
+ row = np.array([0, 1, 2, 3, 4, 5, 6, 7, 8, 9])
301
+ col = np.array([1, 0, 3, 2, 5, 4, 7, 6, 9, 8])
302
+ data_arr = np.ones(10, dtype=np.float32)
303
+ adj_coo = sp.coo_matrix((data_arr, (row, col)), shape=(num_nodes, num_nodes))
304
+ sp.save_npz(osp.join(raw_dir, "reddit_graph.npz"), adj_coo)
305
+
306
+ dataset = datasets.Reddit(root=tmp_dir)
307
+ assert len(dataset) == 1
308
+ assert dataset.num_features == 602
309
+ assert dataset.num_classes == 41
310
+
311
+ data = dataset[0]
312
+ assert data.num_nodes == 10
313
+ assert data.x.shape == (10, 602)
314
+ assert data.y.shape == (10,)
315
+ assert data.edge_index.shape[0] == 2
316
+ assert data.train_mask.shape == (10,)
317
+ assert data.val_mask.shape == (10,)
318
+ assert data.test_mask.shape == (10,)
319
+ assert int(ops.sum(ops.cast(data.train_mask, "int32"))) == 4
320
+ assert int(ops.sum(ops.cast(data.val_mask, "int32"))) == 3
321
+ assert int(ops.sum(ops.cast(data.test_mask, "int32"))) == 3
322
+
@@ -0,0 +1,131 @@
1
+ import os
2
+ import os.path as osp
3
+ import pickle
4
+ from typing import Callable, List, Optional
5
+ from keras import ops
6
+
7
+ from k3_node.data import Data, InMemoryDataset
8
+ from k3_node.io import fs, read_tu_data
9
+
10
+
11
+ class TUDataset(InMemoryDataset):
12
+ r"""A variety of graph kernel benchmark datasets, e.g., "IMDB-BINARY",
13
+ "REDDIT-BINARY" or "PROTEINS", collected from the TU Dortmund University.
14
+
15
+ Args:
16
+ root (str): Root directory where the dataset should be saved.
17
+ name (str): The name of the dataset.
18
+ transform (callable, optional): Transform function for Data objects.
19
+ pre_transform (callable, optional): Pre-transform function.
20
+ pre_filter (callable, optional): Pre-filter function.
21
+ force_reload (bool, optional): Whether to re-process the dataset.
22
+ use_node_attr (bool, optional): Whether to include continuous node attributes.
23
+ use_edge_attr (bool, optional): Whether to include continuous edge attributes.
24
+ cleaned (bool, optional): Whether to use cleaned dataset version.
25
+ """
26
+
27
+ url = "https://www.chrsmrrs.com/graphkerneldatasets"
28
+ cleaned_url = "https://raw.githubusercontent.com/nd7141/graph_datasets/master/datasets"
29
+
30
+ def __init__(
31
+ self,
32
+ root: str,
33
+ name: str,
34
+ transform: Optional[Callable] = None,
35
+ pre_transform: Optional[Callable] = None,
36
+ pre_filter: Optional[Callable] = None,
37
+ force_reload: bool = False,
38
+ use_node_attr: bool = False,
39
+ use_edge_attr: bool = False,
40
+ cleaned: bool = False,
41
+ ):
42
+ self.name = name
43
+ self.cleaned = cleaned
44
+ super().__init__(root, transform, pre_transform, pre_filter, force_reload=force_reload)
45
+
46
+ self.load(self.processed_paths[0])
47
+
48
+ if self._data.x is not None and not use_node_attr:
49
+ num_node_attributes = self.num_node_attributes
50
+ if num_node_attributes > 0:
51
+ self._data.x = self._data.x[:, num_node_attributes:]
52
+ if self._data.edge_attr is not None and not use_edge_attr:
53
+ num_edge_attrs = self.num_edge_attributes
54
+ if num_edge_attrs > 0:
55
+ self._data.edge_attr = self._data.edge_attr[:, num_edge_attrs:]
56
+
57
+ @property
58
+ def raw_dir(self) -> str:
59
+ name = f"raw{'_cleaned' if self.cleaned else ''}"
60
+ return osp.join(self.root, self.name, name)
61
+
62
+ @property
63
+ def processed_dir(self) -> str:
64
+ name = f"processed{'_cleaned' if self.cleaned else ''}"
65
+ return osp.join(self.root, self.name, name)
66
+
67
+ @property
68
+ def num_node_labels(self) -> int:
69
+ return self.sizes.get("num_node_labels", 0)
70
+
71
+ @property
72
+ def num_node_attributes(self) -> int:
73
+ return self.sizes.get("num_node_attributes", 0)
74
+
75
+ @property
76
+ def num_edge_labels(self) -> int:
77
+ return self.sizes.get("num_edge_labels", 0)
78
+
79
+ @property
80
+ def num_edge_attributes(self) -> int:
81
+ return self.sizes.get("num_edge_attributes", 0)
82
+
83
+ @property
84
+ def raw_file_names(self) -> List[str]:
85
+ names = ["A", "graph_indicator"]
86
+ return [f"{self.name}_{name}.txt" for name in names]
87
+
88
+ @property
89
+ def processed_file_names(self) -> str:
90
+ return "data.pt"
91
+
92
+ def download(self):
93
+ url = self.cleaned_url if self.cleaned else self.url
94
+ fs.cp(f"{url}/{self.name}.zip", self.raw_dir, extract=True)
95
+ inner_dir = osp.join(self.raw_dir, self.name)
96
+ if osp.isdir(inner_dir):
97
+ for filename in os.listdir(inner_dir):
98
+ src = osp.join(inner_dir, filename)
99
+ dst = osp.join(self.raw_dir, filename)
100
+ fs.cp(src, dst)
101
+ fs.rm(inner_dir)
102
+
103
+ def process(self):
104
+ self.data, self.slices, sizes = read_tu_data(self.raw_dir, self.name)
105
+
106
+ if self.pre_filter is not None or self.pre_transform is not None:
107
+ data_list = [self.get(idx) for idx in range(len(self))]
108
+ if self.pre_filter is not None:
109
+ data_list = [d for d in data_list if self.pre_filter(d)]
110
+ if self.pre_transform is not None:
111
+ data_list = [self.pre_transform(d) for d in data_list]
112
+ self.data, self.slices = self.collate(data_list)
113
+ self._data_list = None
114
+
115
+ os.makedirs(osp.dirname(self.processed_paths[0]), exist_ok=True)
116
+ saved = False
117
+ if self.processed_paths[0].endswith((".pt", ".pth")):
118
+ try:
119
+ import torch
120
+ data_dict = self._data.to_dict() if hasattr(self._data, "to_dict") else dict(self._data)
121
+ torch.save((data_dict, self.slices, sizes), self.processed_paths[0])
122
+ saved = True
123
+ except Exception:
124
+ pass
125
+ if not saved:
126
+ with open(self.processed_paths[0], "wb") as f:
127
+ pickle.dump((self._data, self.slices, sizes), f)
128
+
129
+ def __repr__(self) -> str:
130
+ return f"{self.name}({len(self)})"
131
+
@@ -0,0 +1,66 @@
1
+ import os.path as osp
2
+ from typing import Callable, Optional
3
+ import numpy as np
4
+ from keras import ops
5
+
6
+ from k3_node.data import Data, InMemoryDataset
7
+ from k3_node.io import fs
8
+
9
+
10
+ class Twitch(InMemoryDataset):
11
+ r"""The Twitch Gamer networks."""
12
+
13
+ url = "https://graphmining.ai/datasets/ptg/twitch"
14
+
15
+ def __init__(
16
+ self,
17
+ root: str,
18
+ name: str,
19
+ transform: Optional[Callable] = None,
20
+ pre_transform: Optional[Callable] = None,
21
+ force_reload: bool = False,
22
+ ):
23
+ self.name = name.upper()
24
+ assert self.name in ["DE", "EN", "ES", "FR", "PT", "RU"]
25
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
26
+ self.load(self.processed_paths[0])
27
+
28
+ @property
29
+ def raw_dir(self) -> str:
30
+ return osp.join(self.root, self.name, "raw")
31
+
32
+ @property
33
+ def processed_dir(self) -> str:
34
+ return osp.join(self.root, self.name, "processed")
35
+
36
+ @property
37
+ def raw_file_names(self) -> str:
38
+ return f"{self.name}.npz"
39
+
40
+ @property
41
+ def processed_file_names(self) -> str:
42
+ return "data.pt"
43
+
44
+ def download(self):
45
+ fs.cp(f"{self.url}/{self.name}.npz", self.raw_dir)
46
+
47
+ def process(self):
48
+ data = np.load(self.raw_paths[0], allow_pickle=True)
49
+ x = data["features"].astype(np.float32)
50
+ y = data["target"].astype(np.int64)
51
+ edge_index = data["edges"].astype(np.int64).T
52
+
53
+ data_obj = Data(
54
+ x=ops.convert_to_tensor(x, dtype="float32"),
55
+ y=ops.convert_to_tensor(y, dtype="int64"),
56
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
57
+ )
58
+
59
+ if self.pre_transform is not None:
60
+ data_obj = self.pre_transform(data_obj)
61
+
62
+ self.save([data_obj], self.processed_paths[0])
63
+
64
+ def __repr__(self) -> str:
65
+ return f"{self.__class__.__name__}{self.name}()"
66
+
@@ -0,0 +1,102 @@
1
+ import os.path as osp
2
+ from typing import Callable, List, Optional
3
+ import numpy as np
4
+ from keras import ops
5
+
6
+ from k3_node.data import Data, InMemoryDataset
7
+ from k3_node.io import fs
8
+ from k3_node.utils.graph import coalesce
9
+
10
+
11
+ class WebKB(InMemoryDataset):
12
+ r"""The WebKB datasets used in the "Geom-GCN: Geometric Graph Convolutional Networks" paper.
13
+ Nodes represent web pages and edges represent hyperlinks between them.
14
+
15
+ Args:
16
+ root (str): Root directory where the dataset should be saved.
17
+ name (str): The name of the dataset ("Cornell", "Texas", "Wisconsin").
18
+ transform (callable, optional): Transform function.
19
+ pre_transform (callable, optional): Pre-transform function.
20
+ force_reload (bool, optional): Whether to re-process the dataset.
21
+ """
22
+
23
+ url = "https://raw.githubusercontent.com/graphdml-uiuc-jlu/geom-gcn/master"
24
+
25
+ def __init__(
26
+ self,
27
+ root: str,
28
+ name: str,
29
+ transform: Optional[Callable] = None,
30
+ pre_transform: Optional[Callable] = None,
31
+ force_reload: bool = False,
32
+ ):
33
+ self.name = name.lower()
34
+ assert self.name in ["cornell", "texas", "wisconsin"]
35
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
36
+ self.load(self.processed_paths[0])
37
+
38
+ @property
39
+ def raw_dir(self) -> str:
40
+ return osp.join(self.root, self.name, "raw")
41
+
42
+ @property
43
+ def processed_dir(self) -> str:
44
+ return osp.join(self.root, self.name, "processed")
45
+
46
+ @property
47
+ def raw_file_names(self) -> List[str]:
48
+ out = ["out1_node_feature_label.txt", "out1_graph_edges.txt"]
49
+ out += [f"{self.name}_split_0.6_0.2_{i}.npz" for i in range(10)]
50
+ return out
51
+
52
+ @property
53
+ def processed_file_names(self) -> str:
54
+ return "data.pt"
55
+
56
+ def download(self):
57
+ for f in self.raw_file_names[:2]:
58
+ fs.cp(f"{self.url}/new_data/{self.name}/{f}", self.raw_dir)
59
+ for f in self.raw_file_names[2:]:
60
+ fs.cp(f"{self.url}/splits/{f}", self.raw_dir)
61
+
62
+ def process(self):
63
+ with open(self.raw_paths[0], "r") as f:
64
+ lines = f.read().split("\n")[1:-1]
65
+ xs = [[float(v) for v in r.split("\t")[1].split(",")] for r in lines]
66
+ ys = [int(r.split("\t")[2]) for r in lines]
67
+
68
+ x = np.array(xs, dtype=np.float32)
69
+ y = np.array(ys, dtype=np.int64)
70
+
71
+ with open(self.raw_paths[1], "r") as f:
72
+ lines = f.read().split("\n")[1:-1]
73
+ edges = [[int(v) for v in r.split("\t")] for r in lines]
74
+ edge_index = np.array(edges, dtype=np.int64).T
75
+ edge_index_t, _ = coalesce(edge_index, num_nodes=x.shape[0])
76
+ edge_index = ops.convert_to_numpy(edge_index_t)
77
+
78
+ train_masks, val_masks, test_masks = [], [], []
79
+ for filepath in self.raw_paths[2:]:
80
+ masks = np.load(filepath)
81
+ train_masks.append(masks["train_mask"])
82
+ val_masks.append(masks["val_mask"])
83
+ test_masks.append(masks["test_mask"])
84
+
85
+ train_mask = np.stack(train_masks, axis=1)
86
+ val_mask = np.stack(val_masks, axis=1)
87
+ test_mask = np.stack(test_masks, axis=1)
88
+
89
+ data = Data(
90
+ x=ops.convert_to_tensor(x, dtype="float32"),
91
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
92
+ y=ops.convert_to_tensor(y, dtype="int64"),
93
+ train_mask=ops.convert_to_tensor(train_mask, dtype="bool"),
94
+ val_mask=ops.convert_to_tensor(val_mask, dtype="bool"),
95
+ test_mask=ops.convert_to_tensor(test_mask, dtype="bool"),
96
+ )
97
+
98
+ if self.pre_transform is not None:
99
+ data = self.pre_transform(data)
100
+
101
+ self.save([data], self.processed_paths[0])
102
+