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,106 @@
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 DBLP(InMemoryDataset):
13
+ r"""A subset of the DBLP computer science bibliography website containing
14
+ four types of entities: authors, papers, terms, and conferences.
15
+ """
16
+
17
+ url = "https://www.dropbox.com/s/yh4grpeks87ugr2/DBLP_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.npy",
36
+ "labels.npy",
37
+ "node_types.npy",
38
+ "train_val_test_idx.npz",
39
+ ]
40
+
41
+ @property
42
+ def processed_file_names(self) -> str:
43
+ return "data.pt"
44
+
45
+ def download(self):
46
+ zip_path = osp.join(self.raw_dir, "DBLP_processed.zip")
47
+ fs.cp(self.url, zip_path, extract=True)
48
+ if osp.exists(zip_path):
49
+ fs.rm(zip_path)
50
+
51
+ def process(self):
52
+ import scipy.sparse as sp
53
+
54
+ data = HeteroData()
55
+ node_types = ["author", "paper", "term", "conference"]
56
+
57
+ for i, node_type in enumerate(node_types[:2]):
58
+ feat_path = osp.join(self.raw_dir, f"features_{i}.npz")
59
+ x = sp.load_npz(feat_path)
60
+ data[node_type].x = ops.convert_to_tensor(
61
+ np.array(x.todense(), dtype=np.float32), dtype="float32"
62
+ )
63
+
64
+ x2 = np.load(osp.join(self.raw_dir, "features_2.npy"))
65
+ data["term"].x = ops.convert_to_tensor(x2.astype(np.float32), dtype="float32")
66
+
67
+ node_type_idx = np.load(osp.join(self.raw_dir, "node_types.npy"))
68
+ data["conference"].num_nodes = int((node_type_idx == 3).sum())
69
+
70
+ y = np.load(osp.join(self.raw_dir, "labels.npy"))
71
+ data["author"].y = ops.convert_to_tensor(y.astype(np.int64), dtype="int64")
72
+
73
+ split = np.load(osp.join(self.raw_dir, "train_val_test_idx.npz"))
74
+ for name in ["train", "val", "test"]:
75
+ idx = split[f"{name}_idx"]
76
+ mask = np.zeros(data["author"].num_nodes, dtype=bool)
77
+ mask[idx] = True
78
+ data["author"][f"{name}_mask"] = ops.convert_to_tensor(mask, dtype="bool")
79
+
80
+ s = {}
81
+ N_a = data["author"].num_nodes
82
+ N_p = data["paper"].num_nodes
83
+ N_t = data["term"].num_nodes
84
+ N_c = data["conference"].num_nodes
85
+ s["author"] = (0, N_a)
86
+ s["paper"] = (N_a, N_a + N_p)
87
+ s["term"] = (N_a + N_p, N_a + N_p + N_t)
88
+ s["conference"] = (N_a + N_p + N_t, N_a + N_p + N_t + N_c)
89
+
90
+ A = sp.load_npz(osp.join(self.raw_dir, "adjM.npz"))
91
+ for src, dst in product(node_types, node_types):
92
+ A_sub = A[s[src][0] : s[src][1], s[dst][0] : s[dst][1]].tocoo()
93
+ if A_sub.nnz > 0:
94
+ row = np.array(A_sub.row, dtype=np.int64)
95
+ col = np.array(A_sub.col, dtype=np.int64)
96
+ edge_index = np.stack([row, col], axis=0)
97
+ data[src, dst].edge_index = ops.convert_to_tensor(edge_index, dtype="int64")
98
+
99
+ if self.pre_transform is not None:
100
+ data = self.pre_transform(data)
101
+
102
+ self.save([data], self.processed_paths[0])
103
+
104
+ def __repr__(self) -> str:
105
+ return f"{self.__class__.__name__}()"
106
+
@@ -0,0 +1,63 @@
1
+ from typing import Callable, Optional
2
+ import numpy as np
3
+
4
+ from k3_node.data import Data, InMemoryDataset
5
+
6
+
7
+ class Digits(InMemoryDataset):
8
+ r"""The handwritten digits of scikit-learn (1,797 images of 8x8 pixels) as graphs.
9
+
10
+ A small, offline stand-in for PyG's ``MNISTSuperpixels``. Every non-blank pixel becomes a node
11
+ with its intensity (in ``[0, 1]``) as the only feature ``x`` and its 2D position ``pos``, scaled
12
+ to ``[0, 28)`` like MNIST. Neighboring pixels (including diagonals) are connected in both
13
+ directions. ``y`` holds the digit (10 classes).
14
+
15
+ Args:
16
+ train (bool): If ``True``, loads the first 1,500 images, otherwise the remaining 297.
17
+ (default: ``True``)
18
+ transform (callable, optional): A function applied to each graph when it is accessed.
19
+
20
+ Example:
21
+ ```python
22
+ from k3_node.datasets import Digits
23
+
24
+ dataset = Digits(train=True)
25
+ print(len(dataset), dataset.num_classes) # 1500 10
26
+ ```
27
+ """
28
+
29
+ num_train = 1500
30
+
31
+ def __init__(self, train: bool = True, transform: Optional[Callable] = None):
32
+ super().__init__(None, transform)
33
+ from sklearn.datasets import load_digits
34
+
35
+ images, labels = load_digits(return_X_y=True)
36
+ images = images.reshape(-1, 8, 8) / 16.0
37
+ split = slice(0, self.num_train) if train else slice(self.num_train, None)
38
+ self.data, self.slices = self.collate(
39
+ [self._to_graph(img, y) for img, y in zip(images[split], labels[split])]
40
+ )
41
+
42
+ @staticmethod
43
+ def _to_graph(image, label):
44
+ rows, cols = np.nonzero(image)
45
+ index = -np.ones((8, 8), dtype=np.int64)
46
+ index[rows, cols] = np.arange(len(rows))
47
+ src, dst = [], []
48
+ for dr in (-1, 0, 1):
49
+ for dc in (-1, 0, 1):
50
+ if dr == 0 and dc == 0:
51
+ continue
52
+ r, c = rows + dr, cols + dc
53
+ ok = (r >= 0) & (r < 8) & (c >= 0) & (c < 8)
54
+ nbr = index[r[ok], c[ok]]
55
+ keep = nbr >= 0
56
+ src.append(np.nonzero(ok)[0][keep])
57
+ dst.append(nbr[keep])
58
+ return Data(
59
+ x=image[rows, cols][:, None].astype(np.float32),
60
+ pos=(np.stack([cols, rows], axis=1) * 3.5 + 1.75).astype(np.float32),
61
+ edge_index=np.stack([np.concatenate(src), np.concatenate(dst)]).astype(np.int64),
62
+ y=np.array([label], dtype=np.int64),
63
+ )
@@ -0,0 +1,60 @@
1
+ import os
2
+ import os.path as osp
3
+ from typing import Callable, List, Optional
4
+ import numpy as np
5
+ from keras import ops
6
+
7
+ from k3_node.data import Data, InMemoryDataset
8
+ from k3_node.io import fs
9
+
10
+
11
+ class EmailEUCore(InMemoryDataset):
12
+ r"""An e-mail communication network of a large European research institution."""
13
+
14
+ urls = [
15
+ "https://snap.stanford.edu/data/email-Eu-core.txt.gz",
16
+ "https://snap.stanford.edu/data/email-Eu-core-department-labels.txt.gz",
17
+ ]
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 ["email-Eu-core.txt", "email-Eu-core-department-labels.txt"]
32
+
33
+ @property
34
+ def processed_file_names(self) -> str:
35
+ return "data.pt"
36
+
37
+ def download(self):
38
+ for url in self.urls:
39
+ filename = osp.basename(url)
40
+ gz_path = osp.join(self.raw_dir, filename)
41
+ fs.cp(url, gz_path, extract=True)
42
+ if osp.exists(gz_path):
43
+ fs.rm(gz_path)
44
+
45
+ def process(self):
46
+ edge_index = np.loadtxt(self.raw_paths[0], dtype=np.int64).T
47
+ labels = np.loadtxt(self.raw_paths[1], dtype=np.int64)
48
+ y = labels[:, 1]
49
+
50
+ data = Data(
51
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
52
+ y=ops.convert_to_tensor(y, dtype="int64"),
53
+ num_nodes=int(y.shape[0]),
54
+ )
55
+
56
+ if self.pre_transform is not None:
57
+ data = self.pre_transform(data)
58
+
59
+ self.save([data], self.processed_paths[0])
60
+
@@ -0,0 +1,158 @@
1
+ import logging
2
+ import os
3
+ import os.path as osp
4
+ from collections import Counter
5
+ from typing import Any, Callable, List, Optional
6
+ import numpy as np
7
+ from keras import ops
8
+
9
+ from k3_node.data import Data, HeteroData, InMemoryDataset
10
+ from k3_node.io import fs
11
+
12
+
13
+ class Entities(InMemoryDataset):
14
+ r"""The relational entities networks "AIFB", "MUTAG", "BGS" and "AM"."""
15
+
16
+ url = "https://data.dgl.ai/dataset/{}.tgz"
17
+
18
+ def __init__(
19
+ self,
20
+ root: str,
21
+ name: str,
22
+ hetero: bool = False,
23
+ transform: Optional[Callable] = None,
24
+ pre_transform: Optional[Callable] = None,
25
+ force_reload: bool = False,
26
+ ):
27
+ self.name = name.lower()
28
+ self.hetero = hetero
29
+ assert self.name in ["aifb", "am", "mutag", "bgs"]
30
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
31
+ self.load(self.processed_paths[0])
32
+
33
+ @property
34
+ def raw_dir(self) -> str:
35
+ return osp.join(self.root, self.name, "raw")
36
+
37
+ @property
38
+ def processed_dir(self) -> str:
39
+ return osp.join(self.root, self.name, "processed")
40
+
41
+ @property
42
+ def num_relations(self) -> int:
43
+ return int(ops.convert_to_numpy(self._data.edge_type).max()) + 1
44
+
45
+ @property
46
+ def num_classes(self) -> int:
47
+ return int(ops.convert_to_numpy(self._data.train_y).max()) + 1
48
+
49
+ @property
50
+ def raw_file_names(self) -> List[str]:
51
+ return [
52
+ f"{self.name}_stripped.nt.gz",
53
+ "completeDataset.tsv",
54
+ "trainingSet.tsv",
55
+ "testSet.tsv",
56
+ ]
57
+
58
+ @property
59
+ def processed_file_names(self) -> str:
60
+ return "hetero_data.pt" if self.hetero else "data.pt"
61
+
62
+ def download(self):
63
+ tgz_path = osp.join(self.raw_dir, f"{self.name}.tgz") # extracted next to it, into raw_dir
64
+ fs.cp(self.url.format(self.name), tgz_path, extract=True)
65
+ if osp.exists(tgz_path):
66
+ fs.rm(tgz_path)
67
+
68
+ def process(self):
69
+ import gzip
70
+ import rdflib as rdf
71
+
72
+ graph_file, task_file, train_file, test_file = self.raw_paths
73
+
74
+ g = rdf.Graph()
75
+ with gzip.open(graph_file, "rb") as f:
76
+ g.parse(file=f, format="nt")
77
+
78
+ freq = Counter(g.predicates())
79
+ relations = sorted(set(g.predicates()), key=lambda p: -freq.get(p, 0))
80
+ subjects = set(g.subjects())
81
+ objects = set(g.objects())
82
+ nodes = list(subjects.union(objects))
83
+
84
+ N = len(nodes)
85
+ R = 2 * len(relations)
86
+
87
+ relations_dict = {rel: i for i, rel in enumerate(relations)}
88
+ nodes_dict = {str(node): i for i, node in enumerate(nodes)}
89
+
90
+ edges = []
91
+ for s, p, o in g.triples((None, None, None)):
92
+ src, dst = nodes_dict[str(s)], nodes_dict[str(o)]
93
+ rel = relations_dict[p]
94
+ edges.append([src, dst, 2 * rel])
95
+ edges.append([dst, src, 2 * rel + 1])
96
+
97
+ edge = np.array(edges, dtype=np.int64).T
98
+ sort_key = N * R * edge[0] + R * edge[1] + edge[2]
99
+ perm = np.argsort(sort_key)
100
+ edge = edge[:, perm]
101
+
102
+ edge_index, edge_type = edge[:2], edge[2]
103
+
104
+ if self.name == "am":
105
+ label_header = "label_cateogory"
106
+ nodes_header = "proxy"
107
+ elif self.name == "aifb":
108
+ label_header = "label_affiliation"
109
+ nodes_header = "person"
110
+ elif self.name == "mutag":
111
+ label_header = "label_mutagenic"
112
+ nodes_header = "bond"
113
+ elif self.name == "bgs":
114
+ label_header = "label_lithogenesis"
115
+ nodes_header = "rock"
116
+
117
+ import pandas as pd
118
+
119
+ labels_df = pd.read_csv(task_file, sep="\t")
120
+ labels_set = set(labels_df[label_header].values.tolist())
121
+ labels_dict = {lab: i for i, lab in enumerate(list(labels_set))}
122
+
123
+ train_labels_df = pd.read_csv(train_file, sep="\t")
124
+ train_indices, train_labels = [], []
125
+ for nod, lab in zip(train_labels_df[nodes_header].values, train_labels_df[label_header].values):
126
+ train_indices.append(nodes_dict[nod])
127
+ train_labels.append(labels_dict[lab])
128
+
129
+ train_idx = np.array(train_indices, dtype=np.int64)
130
+ train_y = np.array(train_labels, dtype=np.int64)
131
+
132
+ test_labels_df = pd.read_csv(test_file, sep="\t")
133
+ test_indices, test_labels = [], []
134
+ for nod, lab in zip(test_labels_df[nodes_header].values, test_labels_df[label_header].values):
135
+ test_indices.append(nodes_dict[nod])
136
+ test_labels.append(labels_dict[lab])
137
+
138
+ test_idx = np.array(test_indices, dtype=np.int64)
139
+ test_y = np.array(test_labels, dtype=np.int64)
140
+
141
+ data = Data(
142
+ edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
143
+ edge_type=ops.convert_to_tensor(edge_type, dtype="int64"),
144
+ train_idx=ops.convert_to_tensor(train_idx, dtype="int64"),
145
+ train_y=ops.convert_to_tensor(train_y, dtype="int64"),
146
+ test_idx=ops.convert_to_tensor(test_idx, dtype="int64"),
147
+ test_y=ops.convert_to_tensor(test_y, dtype="int64"),
148
+ num_nodes=N,
149
+ )
150
+
151
+ if self.hetero:
152
+ data = data.to_heterogeneous(node_type_names=["v"])
153
+
154
+ self.save([data], self.processed_paths[0])
155
+
156
+ def __repr__(self) -> str:
157
+ return f"{self.name.upper()}{self.__class__.__name__}()"
158
+
@@ -0,0 +1,101 @@
1
+ from typing import Any, Callable, Dict, Optional, Union
2
+ import numpy as np
3
+ from keras import ops
4
+
5
+ from k3_node.data import Data, InMemoryDataset
6
+ from k3_node.datasets.graph_generator import GraphGenerator
7
+ from k3_node.datasets.motif_generator import MotifGenerator
8
+
9
+
10
+ class ExplainerDataset(InMemoryDataset):
11
+ r"""Generates a synthetic dataset for evaluating explainability algorithms,
12
+ as described in the "GNNExplainer: Generating Explanations for Graph Neural Networks" paper.
13
+
14
+ Args:
15
+ graph_generator (GraphGenerator or str): The graph generator to use.
16
+ motif_generator (MotifGenerator or str): The motif generator to use.
17
+ num_motifs (int): The number of motifs to attach to the graph.
18
+ num_graphs (int, optional): The number of graphs to generate. (default: 1)
19
+ graph_generator_kwargs (dict, optional): Keyword arguments for graph generator.
20
+ motif_generator_kwargs (dict, optional): Keyword arguments for motif generator.
21
+ transform (callable, optional): Transform function.
22
+ """
23
+
24
+ def __init__(
25
+ self,
26
+ graph_generator: Union[GraphGenerator, str],
27
+ motif_generator: Union[MotifGenerator, str],
28
+ num_motifs: int,
29
+ num_graphs: int = 1,
30
+ graph_generator_kwargs: Optional[Dict[str, Any]] = None,
31
+ motif_generator_kwargs: Optional[Dict[str, Any]] = None,
32
+ transform: Optional[Callable] = None,
33
+ ):
34
+ super().__init__(root=None, transform=transform)
35
+
36
+ if num_motifs <= 0:
37
+ raise ValueError(f"At least one motif needs to be attached (got {num_motifs})")
38
+
39
+ self.graph_generator = GraphGenerator.resolve(
40
+ graph_generator, **(graph_generator_kwargs or {})
41
+ )
42
+ self.motif_generator = MotifGenerator.resolve(
43
+ motif_generator, **(motif_generator_kwargs or {})
44
+ )
45
+ self.num_motifs = num_motifs
46
+
47
+ data_list = [self.get_graph() for _ in range(num_graphs)]
48
+ self.data, self.slices = self.collate(data_list)
49
+
50
+ def get_graph(self) -> Data:
51
+ data = self.graph_generator()
52
+ edge_index_np = ops.convert_to_numpy(data.edge_index)
53
+ num_nodes = data.num_nodes
54
+ num_edges = edge_index_np.shape[1]
55
+
56
+ edge_indices = [edge_index_np]
57
+ node_masks = [np.zeros(num_nodes, dtype=np.float32)]
58
+ edge_masks = [np.zeros(num_edges, dtype=np.float32)]
59
+ ys = [np.zeros(num_nodes, dtype=np.int64)]
60
+
61
+ connecting_nodes = np.random.permutation(num_nodes)[: self.num_motifs]
62
+ for i in connecting_nodes.tolist():
63
+ motif = self.motif_generator()
64
+ motif_ei = ops.convert_to_numpy(motif.edge_index)
65
+ motif_num_nodes = motif.num_nodes
66
+ motif_num_edges = motif_ei.shape[1]
67
+
68
+ edge_indices.append(motif_ei + num_nodes)
69
+ node_masks.append(np.ones(motif_num_nodes, dtype=np.float32))
70
+ edge_masks.append(np.ones(motif_num_edges, dtype=np.float32))
71
+
72
+ j = int(np.random.randint(0, motif_num_nodes)) + num_nodes
73
+ edge_indices.append(np.array([[i, j], [j, i]], dtype=np.int64))
74
+ edge_masks.append(np.zeros(2, dtype=np.float32))
75
+
76
+ if hasattr(motif, "y") and motif.y is not None:
77
+ motif_y = ops.convert_to_numpy(motif.y)
78
+ if np.min(motif_y) == 0:
79
+ ys.append(motif_y + 1)
80
+ else:
81
+ ys.append(motif_y)
82
+ else:
83
+ ys.append(np.ones(motif_num_nodes, dtype=np.int64))
84
+
85
+ num_nodes += motif_num_nodes
86
+
87
+ return Data(
88
+ edge_index=ops.convert_to_tensor(np.concatenate(edge_indices, axis=1), dtype="int64"),
89
+ y=ops.convert_to_tensor(np.concatenate(ys, axis=0), dtype="int64"),
90
+ edge_mask=ops.convert_to_tensor(np.concatenate(edge_masks, axis=0), dtype="float32"),
91
+ node_mask=ops.convert_to_tensor(np.concatenate(node_masks, axis=0), dtype="float32"),
92
+ )
93
+
94
+ def __repr__(self) -> str:
95
+ return (
96
+ f"{self.__class__.__name__}({len(self)}, "
97
+ f"graph_generator={self.graph_generator}, "
98
+ f"motif_generator={self.motif_generator}, "
99
+ f"num_motifs={self.num_motifs})"
100
+ )
101
+
@@ -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 FacebookPagePage(InMemoryDataset):
10
+ r"""The Facebook Page-Page network dataset."""
11
+
12
+ url = "https://graphmining.ai/datasets/ptg/facebook.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 "facebook.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
+