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,202 @@
1
+ from typing import Any, Callable, Dict, List, Optional, Tuple, Union
2
+
3
+ import numpy as np
4
+
5
+ try:
6
+ import torch
7
+ from torch import Tensor
8
+ except ImportError:
9
+ torch = None
10
+ Tensor = type(None)
11
+
12
+ from k3_node.data import Data, HeteroData
13
+ from k3_node.loader.link_loader import EdgeSamplerInput, LinkLoader
14
+ from k3_node.loader.node_loader import HeteroSamplerOutput, SamplerOutput
15
+ from k3_node.loader.sampler_utils import FastGraph, sample_neighbors_hetero, sample_neighbors_homo
16
+
17
+
18
+ class InternalLinkNeighborSampler:
19
+ r"""Pure Python/NumPy link neighborhood sampler."""
20
+ def __init__(
21
+ self,
22
+ data: Union[Data, HeteroData],
23
+ num_neighbors: Union[List[int], Dict[Tuple[str, str, str], List[int]]],
24
+ replace: bool = False,
25
+ subgraph_type: str = 'directional',
26
+ disjoint: bool = False,
27
+ neg_sampling_ratio: float = 0.0,
28
+ ):
29
+ self.data = data
30
+ self.num_neighbors = num_neighbors
31
+ self.replace = replace
32
+ self.subgraph_type = subgraph_type
33
+ self.disjoint = disjoint
34
+ self.neg_sampling_ratio = neg_sampling_ratio
35
+ self.edge_permutation = None
36
+
37
+ def _cached_graph(self):
38
+ # The CSR index depends only on the graph; build it once instead of per batch.
39
+ if getattr(self, "_graph", None) is None:
40
+ self._graph = FastGraph(self.data.edge_index, num_nodes=self.data.num_nodes)
41
+ return self._graph
42
+
43
+ def sample_from_edges(self, input_data: EdgeSamplerInput) -> Union[SamplerOutput, HeteroSamplerOutput]:
44
+ is_torch = torch is not None and isinstance(input_data.row, Tensor)
45
+
46
+ pos_row = input_data.row
47
+ pos_col = input_data.col
48
+ pos_len = pos_row.size(0) if is_torch else len(pos_row)
49
+
50
+ if self.neg_sampling_ratio > 0:
51
+ num_neg = round(self.neg_sampling_ratio * pos_len)
52
+ num_total_nodes = self.data.num_nodes
53
+ if is_torch:
54
+ # As in PyG, both endpoints of a negative are random nodes
55
+ neg_row = torch.randint(0, num_total_nodes, (num_neg,), dtype=torch.long, device=pos_row.device)
56
+ neg_col = torch.randint(0, num_total_nodes, (num_neg,), dtype=torch.long, device=pos_row.device)
57
+ total_row = torch.cat([pos_row, neg_row], dim=0)
58
+ total_col = torch.cat([pos_col, neg_col], dim=0)
59
+ edge_label = torch.cat([torch.ones(pos_len, dtype=torch.float, device=pos_row.device),
60
+ torch.zeros(num_neg, dtype=torch.float, device=pos_row.device)], dim=0)
61
+ else:
62
+ neg_row = np.random.randint(0, num_total_nodes, size=(num_neg,), dtype=np.int64)
63
+ neg_col = np.random.randint(0, num_total_nodes, size=(num_neg,), dtype=np.int64)
64
+ total_row = np.concatenate([pos_row, neg_row], axis=0)
65
+ total_col = np.concatenate([pos_col, neg_col], axis=0)
66
+ edge_label = np.concatenate([np.ones(pos_len, dtype=np.float32), np.zeros(num_neg, dtype=np.float32)], axis=0)
67
+ else:
68
+ total_row = pos_row
69
+ total_col = pos_col
70
+ edge_label = input_data.label
71
+
72
+ # The sampled subgraph starts with the (sorted, unique) seed nodes, so `edge_label_index`
73
+ # is relabeled to their positions, i.e. to local node ids, as in PyG.
74
+ def unique_inverse(values):
75
+ if is_torch:
76
+ return torch.unique(values, return_inverse=True)
77
+ return np.unique(np.asarray(values), return_inverse=True)
78
+
79
+ def stack(a, b):
80
+ return torch.stack([a, b], dim=0) if is_torch else np.stack([a, b], axis=0)
81
+
82
+ if is_torch:
83
+ seed_nodes, inverse = unique_inverse(torch.cat([total_row, total_col], dim=0))
84
+ else:
85
+ seed_nodes, inverse = unique_inverse(np.concatenate([total_row, total_col], axis=0))
86
+ edge_label_index = stack(inverse[:len(total_row)], inverse[len(total_row):])
87
+
88
+ if isinstance(self.data, Data):
89
+ node, row, col, edge, n_counts, e_counts = sample_neighbors_homo(
90
+ edge_index=self.data.edge_index,
91
+ seed_nodes=seed_nodes,
92
+ num_neighbors=self.num_neighbors,
93
+ num_nodes=self.data.num_nodes,
94
+ graph=self._cached_graph(),
95
+ replace=self.replace,
96
+ subgraph_type=self.subgraph_type,
97
+ disjoint=self.disjoint,
98
+ )
99
+ return SamplerOutput(
100
+ node=node,
101
+ row=row,
102
+ col=col,
103
+ edge=edge,
104
+ num_sampled_nodes=n_counts,
105
+ num_sampled_edges=e_counts,
106
+ metadata=(input_data.input_id, edge_label_index, edge_label),
107
+ )
108
+ elif isinstance(self.data, HeteroData):
109
+ edge_index_dict = {k: self.data[k].edge_index for k in self.data.edge_types}
110
+ input_type = input_data.input_type or self.data.edge_types[0]
111
+ src_type, _, dst_type = input_type
112
+
113
+ seed_dict = {k: None for k in self.data.node_types}
114
+ if src_type == dst_type:
115
+ seed_dict[src_type] = seed_nodes
116
+ else:
117
+ seed_dict[src_type], src_inverse = unique_inverse(total_row)
118
+ seed_dict[dst_type], dst_inverse = unique_inverse(total_col)
119
+ edge_label_index = stack(src_inverse, dst_inverse)
120
+
121
+ norm_num_neighbors = self.num_neighbors
122
+ if isinstance(norm_num_neighbors, dict):
123
+ norm_num_neighbors = {self.data._to_canonical(*k): v for k, v in norm_num_neighbors.items()}
124
+
125
+ node_dict, row_dict, col_dict, edge_dict, n_counts, e_counts = sample_neighbors_hetero(
126
+ edge_index_dict=edge_index_dict,
127
+ seed_nodes_dict=seed_dict,
128
+ num_neighbors=norm_num_neighbors,
129
+ replace=self.replace,
130
+ subgraph_type=self.subgraph_type,
131
+ )
132
+ return HeteroSamplerOutput(
133
+ node=node_dict,
134
+ row=row_dict,
135
+ col=col_dict,
136
+ edge=edge_dict,
137
+ num_sampled_nodes=n_counts,
138
+ num_sampled_edges=e_counts,
139
+ metadata=(input_data.input_id, edge_label_index, edge_label),
140
+ )
141
+
142
+ raise TypeError(f"Invalid data type: {type(self.data)}")
143
+
144
+
145
+ class LinkNeighborLoader(LinkLoader):
146
+ r"""A link-based data loader derived as an extension of NeighborLoader."""
147
+ def __init__(
148
+ self,
149
+ data: Union[Data, HeteroData],
150
+ num_neighbors: Union[List[int], Dict[Tuple[str, str, str], List[int]]],
151
+ edge_label_index: Any = None,
152
+ edge_label: Optional[Any] = None,
153
+ edge_label_time: Optional[Any] = None,
154
+ replace: bool = False,
155
+ subgraph_type: str = 'directional',
156
+ disjoint: bool = False,
157
+ temporal_strategy: str = 'uniform',
158
+ neg_sampling: Optional[Any] = None,
159
+ neg_sampling_ratio: Optional[Union[int, float]] = None,
160
+ time_attr: Optional[str] = None,
161
+ weight_attr: Optional[str] = None,
162
+ transform: Optional[Callable] = None,
163
+ transform_sampler_output: Optional[Callable] = None,
164
+ is_sorted: bool = False,
165
+ filter_per_worker: Optional[bool] = None,
166
+ neighbor_sampler: Optional[Any] = None,
167
+ directed: bool = True,
168
+ **kwargs,
169
+ ):
170
+ if not directed:
171
+ subgraph_type = 'induced'
172
+
173
+ ratio = 0.0
174
+ if neg_sampling_ratio is not None:
175
+ ratio = float(neg_sampling_ratio)
176
+ elif neg_sampling is not None:
177
+ ratio = float(getattr(neg_sampling, 'amount', 1.0)) if hasattr(neg_sampling, 'amount') else 1.0
178
+
179
+ if neighbor_sampler is None:
180
+ neighbor_sampler = InternalLinkNeighborSampler(
181
+ data,
182
+ num_neighbors=num_neighbors,
183
+ replace=replace,
184
+ subgraph_type=str(getattr(subgraph_type, 'value', subgraph_type)),
185
+ disjoint=disjoint,
186
+ neg_sampling_ratio=ratio,
187
+ )
188
+
189
+ super().__init__(
190
+ data=data,
191
+ link_sampler=neighbor_sampler,
192
+ edge_label_index=edge_label_index,
193
+ edge_label=edge_label,
194
+ edge_label_time=edge_label_time,
195
+ neg_sampling=neg_sampling,
196
+ neg_sampling_ratio=neg_sampling_ratio,
197
+ transform=transform,
198
+ transform_sampler_output=transform_sampler_output,
199
+ filter_per_worker=filter_per_worker,
200
+ **kwargs,
201
+ )
202
+
@@ -0,0 +1,190 @@
1
+ import glob
2
+ import logging
3
+ import os
4
+ import os.path as osp
5
+ import warnings
6
+ from contextlib import contextmanager
7
+ from typing import Any, Callable, Dict, List, Optional, Union
8
+
9
+ try:
10
+ import psutil
11
+ except ImportError:
12
+ psutil = None
13
+
14
+ try:
15
+ import torch
16
+ except ImportError:
17
+ torch = None
18
+
19
+
20
+ def get_numa_nodes_cores() -> Dict[str, Any]:
21
+ """Parses numa nodes information into a dictionary."""
22
+ numa_node_paths = glob.glob('/sys/devices/system/node/node[0-9]*')
23
+ if not numa_node_paths:
24
+ return {}
25
+
26
+ nodes = {}
27
+ try:
28
+ for node_path in numa_node_paths:
29
+ numa_node_id = int(osp.basename(node_path)[4:])
30
+ thread_siblings = {}
31
+ for cpu_dir in glob.glob(osp.join(node_path, 'cpu[0-9]*')):
32
+ cpu_id = int(osp.basename(cpu_dir)[3:])
33
+ if cpu_id > 0:
34
+ with open(osp.join(cpu_dir, 'online')) as core_online_file:
35
+ core_online = int(core_online_file.read().splitlines()[0])
36
+ else:
37
+ core_online = 1 # cpu0 is always online
38
+ if core_online == 1:
39
+ with open(osp.join(cpu_dir, 'topology', 'core_id')) as core_id_file:
40
+ core_id = int(core_id_file.read().strip())
41
+ if core_id in thread_siblings:
42
+ thread_siblings[core_id].append(cpu_id)
43
+ else:
44
+ thread_siblings[core_id] = [cpu_id]
45
+
46
+ nodes[numa_node_id] = sorted([(k, sorted(v)) for k, v in thread_siblings.items()])
47
+ except (OSError, ValueError, IndexError):
48
+ warnings.warn('Failed to read NUMA info')
49
+ return {}
50
+
51
+ return nodes
52
+
53
+
54
+ class WorkerInitWrapper:
55
+ r"""Wraps the :attr:`worker_init_fn` argument for DataLoader workers."""
56
+ def __init__(self, func: Optional[Callable]) -> None:
57
+ self.func = func
58
+
59
+ def __call__(self, worker_id: int) -> None:
60
+ if self.func is not None:
61
+ self.func(worker_id)
62
+
63
+
64
+ class LogMemoryMixin:
65
+ r"""A context manager to enable logging of memory consumption in
66
+ DataLoader workers.
67
+ """
68
+ def _mem_init_fn(self, worker_id: int) -> None:
69
+ if psutil is not None:
70
+ proc = psutil.Process(os.getpid())
71
+ memory = proc.memory_info().rss / (1024 * 1024)
72
+ logging.debug(f"Worker {worker_id} @ PID {proc.pid}: {memory:.2f} MB")
73
+ self._old_worker_init_fn(worker_id)
74
+
75
+ @contextmanager
76
+ def enable_memory_log(self):
77
+ self._old_worker_init_fn = WorkerInitWrapper(getattr(self, 'worker_init_fn', None))
78
+ try:
79
+ self.worker_init_fn = self._mem_init_fn
80
+ yield
81
+ finally:
82
+ self.worker_init_fn = self._old_worker_init_fn
83
+
84
+
85
+ class MultithreadingMixin:
86
+ r"""A context manager to enable multi-threading in DataLoader workers."""
87
+ def _mt_init_fn(self, worker_id: int) -> None:
88
+ if torch is not None:
89
+ try:
90
+ torch.set_num_threads(int(self._worker_threads))
91
+ except IndexError as e:
92
+ raise ValueError(f"Cannot set {self._worker_threads} threads in worker {worker_id}") from e
93
+ self._old_worker_init_fn(worker_id)
94
+
95
+ @contextmanager
96
+ def enable_multithreading(self, worker_threads: Optional[int] = None):
97
+ num_workers = getattr(self, 'num_workers', 0)
98
+ if not num_workers > 0:
99
+ raise ValueError(f"'enable_multithreading' needs to be performed with at least one worker (got {num_workers})")
100
+
101
+ if torch is not None:
102
+ if worker_threads is None:
103
+ worker_threads = torch.get_num_threads() // num_workers
104
+ if worker_threads > torch.get_num_threads():
105
+ raise ValueError(
106
+ f"'worker_threads' should be smaller than total available threads {torch.get_num_threads()} (got {worker_threads})"
107
+ )
108
+ context = torch.multiprocessing.get_context()._name
109
+ if context != 'spawn':
110
+ raise ValueError(f"'enable_multithreading' can only be used with 'spawn' multiprocessing context (got {context})")
111
+ else:
112
+ if worker_threads is None:
113
+ worker_threads = 1
114
+
115
+ self._worker_threads = worker_threads
116
+ self._old_worker_init_fn = WorkerInitWrapper(getattr(self, 'worker_init_fn', None))
117
+ try:
118
+ logging.debug(f"Using {worker_threads} threads in each worker")
119
+ self.worker_init_fn = self._mt_init_fn
120
+ yield
121
+ finally:
122
+ self.worker_init_fn = self._old_worker_init_fn
123
+
124
+
125
+ class AffinityMixin:
126
+ r"""A context manager to enable CPU affinity for data loader workers."""
127
+ def _aff_init_fn(self, worker_id: int) -> None:
128
+ try:
129
+ worker_cores = self.loader_cores[worker_id]
130
+ if not isinstance(worker_cores, list):
131
+ worker_cores = [worker_cores]
132
+
133
+ if torch is not None and torch.multiprocessing.get_context()._name == 'spawn':
134
+ torch.set_num_threads(len(worker_cores))
135
+
136
+ if psutil is not None:
137
+ psutil.Process().cpu_affinity(worker_cores)
138
+ except IndexError as e:
139
+ raise ValueError(f"Cannot use CPU affinity for worker ID {worker_id} on CPU {self.loader_cores}") from e
140
+
141
+ self._old_worker_init_fn(worker_id)
142
+
143
+ @contextmanager
144
+ def enable_cpu_affinity(self, loader_cores: Optional[Union[List[List[int]], List[int]]] = None):
145
+ num_workers = getattr(self, 'num_workers', 0)
146
+ if not num_workers > 0:
147
+ raise ValueError(f"'enable_cpu_affinity' should be used with at least one worker (got {num_workers})")
148
+ if loader_cores and len(loader_cores) != num_workers:
149
+ raise ValueError(
150
+ f"The number of loader cores ({len(loader_cores)}) in 'enable_cpu_affinity' should match number of workers ({num_workers})"
151
+ )
152
+
153
+ from k3_node.data import HeteroData
154
+ if hasattr(self, 'data') and isinstance(self.data, HeteroData):
155
+ warnings.warn(
156
+ "Due to conflicting parallelization methods it is not advised to use affinitization with 'HeteroData' datasets.",
157
+ stacklevel=2,
158
+ )
159
+
160
+ self.loader_cores = loader_cores[:] if loader_cores else None
161
+ if self.loader_cores is None:
162
+ numa_info = get_numa_nodes_cores()
163
+ if numa_info and len(numa_info.get(0, [])) > num_workers:
164
+ node0_cores = [cpus[0] for core_id, cpus in numa_info[0]]
165
+ node0_cores.sort()
166
+ elif psutil is not None:
167
+ node0_cores = list(range(psutil.cpu_count(logical=False) or 1))
168
+ else:
169
+ node0_cores = list(range(os.cpu_count() or 1))
170
+
171
+ if len(node0_cores) < num_workers:
172
+ raise ValueError(f"More workers ({num_workers}) than available cores ({len(node0_cores)})")
173
+
174
+ if torch is not None and torch.multiprocessing.get_context()._name == 'spawn':
175
+ work_thread_pool = int(len(node0_cores) / num_workers)
176
+ self.loader_cores = [
177
+ list(range(work_thread_pool * i, work_thread_pool * (i + 1)))
178
+ for i in range(num_workers)
179
+ ]
180
+ else:
181
+ self.loader_cores = node0_cores[:num_workers]
182
+
183
+ self._old_worker_init_fn = WorkerInitWrapper(getattr(self, 'worker_init_fn', None))
184
+ try:
185
+ logging.debug(f"{num_workers} data loader workers assigned to CPUs {self.loader_cores}")
186
+ self.worker_init_fn = self._aff_init_fn
187
+ yield
188
+ finally:
189
+ self.worker_init_fn = self._old_worker_init_fn
190
+
@@ -0,0 +1,159 @@
1
+ from typing import Any, Callable, Dict, List, Optional, Tuple, Union
2
+
3
+ from k3_node.data import Data, HeteroData
4
+ from k3_node.loader.node_loader import HeteroSamplerOutput, NodeLoader, NodeSamplerInput, SamplerOutput
5
+ import numpy as np
6
+
7
+ from k3_node.loader.sampler_utils import (FastGraph, sample_neighbors_disjoint, sample_neighbors_hetero,
8
+ sample_neighbors_homo)
9
+
10
+
11
+ class InternalNeighborSampler:
12
+ r"""Pure Python/NumPy neighborhood sampling engine."""
13
+ def __init__(
14
+ self,
15
+ data: Union[Data, HeteroData],
16
+ num_neighbors: Union[List[int], Dict[Tuple[str, str, str], List[int]]],
17
+ replace: bool = False,
18
+ subgraph_type: str = 'directional',
19
+ disjoint: bool = False,
20
+ ):
21
+ self.data = data
22
+ self.num_neighbors = num_neighbors
23
+ self.replace = replace
24
+ self.subgraph_type = subgraph_type
25
+ self.disjoint = disjoint
26
+ self.edge_permutation = None
27
+ self._graph = None # CSR index, built once and reused for every batch
28
+
29
+ def sample_from_nodes(self, input_data: NodeSamplerInput) -> Union[SamplerOutput, HeteroSamplerOutput]:
30
+ if isinstance(self.data, Data):
31
+ if self._graph is None:
32
+ self._graph = FastGraph(self.data.edge_index, num_nodes=self.data.num_nodes)
33
+ if self.disjoint: # one separate subgraph per seed node
34
+ seeds = input_data.node
35
+ is_torch = hasattr(seeds, "detach")
36
+ out = sample_neighbors_disjoint(self._graph, seeds.detach().cpu().numpy() if is_torch else seeds,
37
+ self.num_neighbors, replace=self.replace)
38
+ node, row, col, edge, batch, n_counts, e_counts = out
39
+ if is_torch:
40
+ import torch
41
+ node, row, col, edge, batch = (torch.from_numpy(a) for a in (node, row, col, edge, batch))
42
+ return SamplerOutput(node=node, row=row, col=col, edge=edge, batch=batch,
43
+ num_sampled_nodes=n_counts, num_sampled_edges=e_counts,
44
+ metadata=(input_data.input_id, input_data.time))
45
+ node, row, col, edge, n_counts, e_counts = sample_neighbors_homo(
46
+ edge_index=self.data.edge_index,
47
+ seed_nodes=input_data.node,
48
+ num_neighbors=self.num_neighbors,
49
+ num_nodes=self.data.num_nodes,
50
+ replace=self.replace,
51
+ subgraph_type=self.subgraph_type,
52
+ disjoint=self.disjoint,
53
+ graph=self._graph,
54
+ )
55
+ return SamplerOutput(
56
+ node=node,
57
+ row=row,
58
+ col=col,
59
+ edge=edge,
60
+ num_sampled_nodes=n_counts,
61
+ num_sampled_edges=e_counts,
62
+ metadata=(input_data.input_id, input_data.time),
63
+ )
64
+ elif isinstance(self.data, HeteroData):
65
+ edge_index_dict = {}
66
+ for edge_type in self.data.edge_types:
67
+ canonical = self.data._to_canonical(*edge_type) if hasattr(self.data, '_to_canonical') else edge_type
68
+ edge_index_dict[canonical] = self.data[edge_type].edge_index
69
+
70
+ node_type = input_data.input_type or self.data.node_types[0]
71
+ seed_dict = {k: None for k in self.data.node_types}
72
+ seed_dict[node_type] = input_data.node
73
+
74
+ # Normalize num_neighbors for hetero
75
+ if isinstance(self.num_neighbors, dict):
76
+ norm_num_neighbors = {}
77
+ for k, v in self.num_neighbors.items():
78
+ can = self.data._to_canonical(*k) if hasattr(self.data, '_to_canonical') else k
79
+ norm_num_neighbors[can] = v
80
+ else:
81
+ norm_num_neighbors = self.num_neighbors
82
+
83
+ node_dict, row_dict, col_dict, edge_dict, n_counts, e_counts = sample_neighbors_hetero(
84
+ edge_index_dict=edge_index_dict,
85
+ seed_nodes_dict=seed_dict,
86
+ num_neighbors=norm_num_neighbors,
87
+ replace=self.replace,
88
+ subgraph_type=self.subgraph_type,
89
+ )
90
+ return HeteroSamplerOutput(
91
+ node=node_dict,
92
+ row=row_dict,
93
+ col=col_dict,
94
+ edge=edge_dict,
95
+ num_sampled_nodes=n_counts,
96
+ num_sampled_edges=e_counts,
97
+ metadata=(input_data.input_id, input_data.time),
98
+ )
99
+
100
+ raise TypeError(f"Invalid data type for sampling: {type(self.data)}")
101
+
102
+
103
+ class NeighborLoader(NodeLoader):
104
+ r"""A data loader that performs neighbor sampling as introduced in
105
+ "Inductive Representation Learning on Large Graphs".
106
+
107
+ Args:
108
+ data (Data or HeteroData): The graph data object.
109
+ num_neighbors (List[int] or Dict[EdgeType, List[int]]): Number of neighbors to sample per iteration.
110
+ input_nodes (Tensor or str or Tuple[str, Tensor], optional): Seed nodes. (default: :obj:`None`)
111
+ replace (bool, optional): Sample with replacement. (default: :obj:`False`)
112
+ subgraph_type (str, optional): :obj:`"directional"`, :obj:`"bidirectional"`, or :obj:`"induced"`.
113
+ (default: :obj:`"directional"`)
114
+ disjoint (bool, optional): If :obj:`True`, creates disjoint subgraphs per seed node. (default: :obj:`False`)
115
+ **kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
116
+ """
117
+ def __init__(
118
+ self,
119
+ data: Union[Data, HeteroData],
120
+ num_neighbors: Union[List[int], Dict[Tuple[str, str, str], List[int]]],
121
+ input_nodes: Any = None,
122
+ input_time: Optional[Any] = None,
123
+ replace: bool = False,
124
+ subgraph_type: str = 'directional',
125
+ disjoint: bool = False,
126
+ temporal_strategy: str = 'uniform',
127
+ time_attr: Optional[str] = None,
128
+ weight_attr: Optional[str] = None,
129
+ transform: Optional[Callable] = None,
130
+ transform_sampler_output: Optional[Callable] = None,
131
+ is_sorted: bool = False,
132
+ filter_per_worker: Optional[bool] = None,
133
+ neighbor_sampler: Optional[Any] = None,
134
+ directed: bool = True,
135
+ **kwargs,
136
+ ):
137
+ if not directed:
138
+ subgraph_type = 'induced'
139
+
140
+ if neighbor_sampler is None:
141
+ neighbor_sampler = InternalNeighborSampler(
142
+ data,
143
+ num_neighbors=num_neighbors,
144
+ replace=replace,
145
+ subgraph_type=str(getattr(subgraph_type, 'value', subgraph_type)),
146
+ disjoint=disjoint,
147
+ )
148
+
149
+ super().__init__(
150
+ data=data,
151
+ node_sampler=neighbor_sampler,
152
+ input_nodes=input_nodes,
153
+ input_time=input_time,
154
+ transform=transform,
155
+ transform_sampler_output=transform_sampler_output,
156
+ filter_per_worker=filter_per_worker,
157
+ **kwargs,
158
+ )
159
+