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,334 @@
1
+ """Makes graphs and graph loaders usable directly with ``model.fit`` / ``evaluate`` / ``predict``.
2
+
3
+ A model receives each batch as a :class:`GraphBatch`: a named tuple with one field per graph
4
+ attribute (``data.x``, ``data.edge_index``, ``data.batch``, ...), so its ``call`` reads like
5
+ PyG's ``forward``. The target (``y`` or ``edge_label``) and sample weights are passed to Keras
6
+ separately.
7
+ """
8
+ import collections
9
+ from typing import Any, Dict, Optional, Sequence, Tuple
10
+
11
+ import numpy as np
12
+ import keras
13
+ from keras import ops
14
+
15
+
16
+ class _GraphBatchMixin:
17
+ __slots__ = ()
18
+
19
+ @property
20
+ def num_graphs(self) -> Optional[int]:
21
+ """Number of graphs in the batch.
22
+
23
+ It is a Python ``int`` whenever the batch shape is known, including inside compiled
24
+ (``jax.jit`` / XLA) training steps, so it can be passed as ``size`` to the global pooling
25
+ functions: ``global_add_pool(x, data.batch, data.num_graphs)``.
26
+ """
27
+ ptr = getattr(self, "ptr", None)
28
+ if ptr is not None:
29
+ n = ptr.shape[0]
30
+ return None if n is None else int(n) - 1
31
+ return 1 if getattr(self, "batch", None) is None else None
32
+
33
+ @property
34
+ def num_nodes(self) -> Optional[int]:
35
+ x = getattr(self, "x", None)
36
+ return None if x is None else x.shape[0]
37
+
38
+
39
+ _TYPES: Dict[Tuple[str, ...], type] = {}
40
+
41
+
42
+ def _graph_batch_type(fields: Tuple[str, ...]) -> type:
43
+ if fields not in _TYPES:
44
+ base = collections.namedtuple("GraphBatch", fields)
45
+ _TYPES[fields] = type("GraphBatch", (base, _GraphBatchMixin), {"__slots__": ()})
46
+ return _TYPES[fields]
47
+
48
+
49
+ def _as_array(value) -> Optional[np.ndarray]:
50
+ if isinstance(value, (bool, int, float, str)) or value is None:
51
+ return None
52
+ try:
53
+ array = np.asarray(ops.convert_to_numpy(value))
54
+ except Exception:
55
+ return None
56
+ if array.dtype.kind not in "biuf" or array.ndim == 0: # skip strings (e.g. SMILES) and objects
57
+ return None
58
+ if array.dtype == np.float64:
59
+ array = array.astype(np.float32)
60
+ elif array.dtype == np.int64:
61
+ array = array.astype(np.int32)
62
+ return array
63
+
64
+
65
+ HeteroGraphBatch = collections.namedtuple("HeteroGraphBatch", ["x_dict", "edge_index_dict"])
66
+
67
+
68
+ def _hetero_keras_batch(data, target, mask, normalize_mask, node_type):
69
+ x_dict = {k: _as_array(v) for k, v in data.collect("x").items()}
70
+ edge_index_dict = {k: _as_array(v) for k, v in data.collect("edge_index").items()}
71
+ inputs = HeteroGraphBatch(x_dict, edge_index_dict)
72
+ if node_type is None:
73
+ return (inputs,)
74
+ store = data[node_type]
75
+ y = _as_array(getattr(store, target or "y"))
76
+ if mask is None:
77
+ return inputs, y
78
+ weight = np.asarray(ops.convert_to_numpy(getattr(store, mask))).astype(np.float32)
79
+ if normalize_mask and weight.sum() > 0:
80
+ weight = weight * (weight.shape[0] / weight.sum())
81
+ return inputs, y, weight
82
+
83
+
84
+ def to_keras_batch(
85
+ data: Any,
86
+ target: Optional[str] = None,
87
+ mask: Optional[str] = None,
88
+ normalize_mask: bool = True,
89
+ index: Optional[str] = None,
90
+ node_type: Optional[str] = None,
91
+ ):
92
+ r"""Converts a :class:`~k3_node.data.Data` / :class:`~k3_node.data.Batch` into Keras' format.
93
+
94
+ For a :class:`~k3_node.data.HeteroData` graph, ``inputs`` is a :class:`HeteroGraphBatch`
95
+ (``x_dict`` and ``edge_index_dict``) and ``target`` / ``mask`` are read from ``node_type``.
96
+
97
+ Returns ``(inputs,)``, ``(inputs, y)`` or ``(inputs, y, sample_weight)``, where ``inputs`` is a
98
+ :class:`GraphBatch` holding every array attribute except the target and ``*_mask`` attributes.
99
+
100
+ Args:
101
+ data: The graph (or mini-batch of graphs).
102
+ target (str, optional): Attribute to predict. Defaults to ``"edge_label"`` if present
103
+ (link prediction), else ``"y"``.
104
+ mask (str, optional): Node mask attribute (e.g. ``"train_mask"``) used as sample weights.
105
+ normalize_mask (bool): Scale sample weights so the loss is the mean over the weighted
106
+ nodes (as in PyG), not Keras' sum divided by the number of all nodes.
107
+ index (str, optional): Node index attribute (e.g. ``"train_idx"``) for datasets that store
108
+ splits as node indices, with ``target`` holding the labels of exactly those nodes
109
+ (e.g. ``"train_y"``). The labels are placed at their nodes and only these nodes count.
110
+ """
111
+ if hasattr(data, "node_types") and hasattr(data, "edge_types"):
112
+ return _hetero_keras_batch(data, target, mask, normalize_mask, node_type)
113
+ attrs = data.to_dict() if hasattr(data, "to_dict") else dict(vars(data))
114
+ if target is None:
115
+ target = "edge_label" if attrs.get("edge_label") is not None else "y"
116
+
117
+ fields = {}
118
+ for key in sorted(attrs):
119
+ if key == target or key.endswith(("_mask", "_idx")):
120
+ continue
121
+ array = _as_array(attrs[key])
122
+ if array is not None:
123
+ fields[key] = array
124
+ for key in ("batch", "ptr"): # Batch exposes these as properties
125
+ value = getattr(data, key, None)
126
+ if key not in fields and value is not None and _as_array(value) is not None:
127
+ fields[key] = _as_array(value)
128
+ if "batch" in fields and fields["batch"].shape[0] == 0:
129
+ # Graphs without nodes (e.g. batches of events): an empty `batch` would come first and
130
+ # make Keras weight the reported loss by 0.
131
+ fields.pop("batch")
132
+ fields.pop("ptr", None)
133
+ inputs = _graph_batch_type(tuple(fields))(**fields)
134
+
135
+ y = _as_array(attrs.get(target))
136
+ if y is None:
137
+ return (inputs,)
138
+
139
+ weight = None
140
+ if index is not None:
141
+ if attrs.get(index) is None:
142
+ raise ValueError(f"The graph has no index attribute '{index}'")
143
+ idx = np.asarray(ops.convert_to_numpy(attrs[index])).astype(np.int64)
144
+ num_nodes = attrs.get("num_nodes") or data.num_nodes
145
+ y_full = np.zeros((num_nodes,) + y.shape[1:], dtype=y.dtype)
146
+ y_full[idx] = y
147
+ y, weight = y_full, np.zeros(num_nodes, dtype=np.float32)
148
+ weight[idx] = 1.0
149
+ elif mask is not None:
150
+ if attrs.get(mask) is None:
151
+ raise ValueError(f"The graph has no mask attribute '{mask}'")
152
+ weight = np.asarray(ops.convert_to_numpy(attrs[mask])).astype(np.float32)
153
+ elif isinstance(attrs.get("batch_size"), int) and y.shape[0] == attrs.get("num_nodes", y.shape[0]):
154
+ # NeighborLoader: only the first `batch_size` (seed) nodes are supervised
155
+ weight = np.zeros(y.shape[0], dtype=np.float32)
156
+ weight[: attrs["batch_size"]] = 1.0
157
+ if weight is None:
158
+ return inputs, y
159
+ if normalize_mask and weight.sum() > 0:
160
+ weight = weight * (weight.shape[0] / weight.sum())
161
+ return inputs, y, weight
162
+
163
+
164
+ class KerasLoaderMixin(keras.utils.PyDataset):
165
+ r"""Lets a graph loader be passed straight to ``model.fit`` / ``evaluate`` / ``predict``.
166
+
167
+ Iterating the loader still yields :class:`~k3_node.data.Batch` objects; Keras instead reads
168
+ batches through ``__getitem__`` in the format produced by :func:`to_keras_batch`.
169
+ """
170
+
171
+ # Class-level defaults stand in for PyDataset.__init__, which loader constructors don't call.
172
+ _workers = 1
173
+ _use_multiprocessing = False
174
+ _max_queue_size = 10
175
+ keras_target: Optional[str] = None
176
+ keras_mask: Optional[str] = None
177
+
178
+ def _keras_index_batches(self):
179
+ if getattr(self, "_keras_batches", None) is None:
180
+ batch_sampler = getattr(self, "batch_sampler", None)
181
+ if batch_sampler is not None and getattr(self, "batch_size", None) is not None:
182
+ self._keras_batches = list(iter(batch_sampler))
183
+ else:
184
+ self._keras_batches = False # no random access: stream from the iterator
185
+ return self._keras_batches
186
+
187
+ def __getitem__(self, index):
188
+ index_batches = self._keras_index_batches()
189
+ if index_batches:
190
+ batch = self.collate_fn([self.dataset[i] for i in index_batches[index]])
191
+ else:
192
+ if index == 0 or getattr(self, "_keras_iter", None) is None:
193
+ self._keras_iter = iter(self)
194
+ try:
195
+ batch = next(self._keras_iter)
196
+ except StopIteration:
197
+ self._keras_iter = iter(self)
198
+ batch = next(self._keras_iter)
199
+ return to_keras_batch(batch, target=self.keras_target, mask=self.keras_mask)
200
+
201
+ def with_mask(self, mask: str):
202
+ r"""Only the nodes in the node mask attribute ``mask`` (e.g. ``"train_mask"``) of each batch
203
+ count in the loss and in ``weighted_metrics``. Returns the loader, for chaining."""
204
+ self.keras_mask = mask
205
+ return self
206
+
207
+ def with_target(self, target: str):
208
+ r"""Sets the attribute Keras predicts (default: ``"edge_label"`` if present, else ``"y"``).
209
+ Returns the loader, for chaining."""
210
+ self.keras_target = target
211
+ return self
212
+
213
+ @property
214
+ def num_batches(self):
215
+ return len(self)
216
+
217
+ def on_epoch_end(self):
218
+ self._keras_batches = None # reshuffle next epoch
219
+ self._keras_iter = None
220
+
221
+
222
+ def _patch_tf_signature():
223
+ """On TensorFlow, Keras fixes every tensor size that is the same in the first few batches it
224
+ reads (except the first axis). A sampling loader with fewer batches than that (e.g. a single
225
+ validation batch) returns different subgraphs on every call, so Keras would fix sizes that
226
+ change later. For K3-Node's loaders, the first batches are therefore read at least twice:
227
+ sizes that vary between calls become variable, and deterministic batches keep static shapes.
228
+ Falls back to Keras' behavior if its internals change."""
229
+ try:
230
+ from keras.src.trainers.data_adapters import data_adapter_utils
231
+ from keras.src.trainers.data_adapters.py_dataset_adapter import PyDatasetAdapter
232
+ except ImportError:
233
+ return
234
+ if getattr(PyDatasetAdapter, "_k3_node_patched", False):
235
+ return
236
+ original = PyDatasetAdapter.get_tf_dataset
237
+
238
+ def get_tf_dataset(self):
239
+ dataset = getattr(self, "py_dataset", None)
240
+ if getattr(self, "_output_signature", "missing") is None and isinstance(dataset, KerasLoaderMixin):
241
+ try:
242
+ num_samples = max(data_adapter_utils.NUM_BATCHES_FOR_TENSOR_SPEC, 2)
243
+ num_batches = dataset.num_batches or num_samples
244
+ # e.g. 3 samples of a 1-batch loader read batch 0 three times
245
+ batches = [self._standardize_batch(dataset[i % num_batches]) for i in range(num_samples)]
246
+ self._output_signature = data_adapter_utils.get_tensor_spec(batches)
247
+ except Exception:
248
+ self._output_signature = None
249
+ return original(self)
250
+
251
+ PyDatasetAdapter.get_tf_dataset = get_tf_dataset
252
+ PyDatasetAdapter._k3_node_patched = True
253
+
254
+
255
+ class FullGraphDataset(keras.utils.PyDataset):
256
+ r"""Feeds one whole graph to ``model.fit`` / ``evaluate`` / ``predict`` as a single batch.
257
+
258
+ Passing node arrays together with ``edge_index`` directly to ``fit`` fails, because Keras
259
+ slices every input along its first axis and ``edge_index`` has shape ``[2, num_edges]``.
260
+ This dataset yields the full graph as one batch per epoch instead. The model receives a
261
+ :class:`GraphBatch` (``data.x``, ``data.edge_index``, ...).
262
+
263
+ Args:
264
+ data: A :class:`~k3_node.data.Data` object.
265
+ mask (str, optional): Node mask attribute (e.g. ``"train_mask"``) whose nodes are used in
266
+ the loss and in ``weighted_metrics``. (default: :obj:`None`, all nodes)
267
+ target (str, optional): Attribute to predict. (default: ``"y"``)
268
+ index (str, optional): For splits stored as node indices: the index attribute (e.g.
269
+ ``"train_idx"``), with ``target`` giving the labels of those nodes (e.g. ``"train_y"``).
270
+ neg_sampling_ratio (float, optional): For link prediction: adds this many random
271
+ non-edges per labeled edge (label 0), sampled anew every epoch.
272
+ node_type (str, optional): For a heterogeneous graph: the node type whose ``target`` is
273
+ predicted and whose ``mask`` selects the nodes. The model receives ``data.x_dict``
274
+ and ``data.edge_index_dict``.
275
+
276
+ Example:
277
+ ```python
278
+ model.compile(optimizer="adam", loss=..., weighted_metrics=["accuracy"])
279
+ model.fit(FullGraphDataset(data, mask="train_mask"), epochs=200)
280
+ model.evaluate(FullGraphDataset(data, mask="test_mask"))
281
+ ```
282
+ """
283
+
284
+ def __init__(self, data, mask: Optional[str] = None, target: Optional[str] = None,
285
+ index: Optional[str] = None, neg_sampling_ratio: Optional[float] = None,
286
+ node_type: Optional[str] = None, **kwargs):
287
+ super().__init__(**kwargs)
288
+ self._data, self._neg_sampling_ratio = data, neg_sampling_ratio
289
+ self._kwargs = dict(target=target, mask=mask, index=index)
290
+ if node_type is not None:
291
+ self._kwargs["node_type"] = node_type
292
+ self._batch = None if neg_sampling_ratio else to_keras_batch(data, **self._kwargs)
293
+
294
+ def __len__(self):
295
+ return 1
296
+
297
+ def __getitem__(self, index):
298
+ if index != 0:
299
+ raise IndexError(index)
300
+ if self._neg_sampling_ratio: # fresh negative edges every epoch
301
+ return to_keras_batch(add_negative_edges(self._data, self._neg_sampling_ratio), **self._kwargs)
302
+ return self._batch
303
+
304
+
305
+ def add_negative_edges(data, ratio: float = 1.0):
306
+ r"""Returns a copy of a link prediction graph with random non-edges added to its labeled edges.
307
+
308
+ ``ratio`` negatives are sampled per labeled edge in ``edge_label_index`` (or per edge of
309
+ ``edge_index`` if there are no labeled edges), avoiding the edges of ``edge_index``. They are
310
+ appended to ``edge_label_index`` with label 0 in ``edge_label``.
311
+ """
312
+ import copy
313
+
314
+ from k3_node.models.utils import negative_sampling
315
+
316
+ pos = getattr(data, "edge_label_index", None)
317
+ pos = np.asarray(ops.convert_to_numpy(data.edge_index if pos is None else pos))
318
+ label = getattr(data, "edge_label", None)
319
+ label = np.ones(pos.shape[1], np.float32) if label is None else np.asarray(ops.convert_to_numpy(label))
320
+ neg = np.asarray(ops.convert_to_numpy(negative_sampling(
321
+ data.edge_index, data.num_nodes, num_neg_samples=int(round(ratio * pos.shape[1])))))
322
+ out = copy.copy(data)
323
+ out.edge_label_index = np.concatenate([pos, neg.astype(pos.dtype)], axis=1)
324
+ out.edge_label = np.concatenate([label, np.zeros(neg.shape[1], label.dtype)])
325
+ return out
326
+
327
+
328
+ def loader_bases(base: type) -> tuple:
329
+ """Base classes for a graph loader: its data-loader base plus :class:`KerasLoaderMixin`."""
330
+ return (base, KerasLoaderMixin) if base is not object else (KerasLoaderMixin,)
331
+
332
+
333
+ if keras.config.backend() == "tensorflow":
334
+ _patch_tf_signature()
@@ -0,0 +1,179 @@
1
+ from dataclasses import dataclass
2
+ from typing import Any, Callable, Dict, Iterator, List, Optional, Tuple, Union
3
+
4
+ import numpy as np
5
+
6
+ try:
7
+ import torch
8
+ from torch import Tensor
9
+ except ImportError:
10
+ torch = None
11
+ Tensor = type(None)
12
+
13
+ from k3_node.data import Data, HeteroData
14
+ from k3_node.loader.base import BaseDataLoader, DataLoaderIterator
15
+ from k3_node.loader.mixin import AffinityMixin, LogMemoryMixin, MultithreadingMixin
16
+ from k3_node.loader.node_loader import HeteroSamplerOutput, SamplerOutput
17
+ from k3_node.loader.utils import (
18
+ filter_data,
19
+ filter_hetero_data,
20
+ get_edge_label_index,
21
+ infer_filter_per_worker,
22
+ )
23
+ from k3_node.loader.keras_dataset import loader_bases
24
+
25
+
26
+ @dataclass
27
+ class EdgeSamplerInput:
28
+ input_id: Optional[Any]
29
+ row: Any
30
+ col: Any
31
+ label: Optional[Any] = None
32
+ time: Optional[Any] = None
33
+ input_type: Optional[Tuple[str, str, str]] = None
34
+
35
+ def __getitem__(self, index: Any) -> 'EdgeSamplerInput':
36
+ if torch is not None and not isinstance(index, Tensor):
37
+ index = torch.as_tensor(index, dtype=torch.long)
38
+ return EdgeSamplerInput(
39
+ input_id=self.input_id[index] if self.input_id is not None else index,
40
+ row=self.row[index],
41
+ col=self.col[index],
42
+ label=self.label[index] if self.label is not None else None,
43
+ time=self.time[index] if self.time is not None else None,
44
+ input_type=self.input_type,
45
+ )
46
+
47
+
48
+ class LinkLoader(*loader_bases(BaseDataLoader), AffinityMixin, MultithreadingMixin, LogMemoryMixin):
49
+ r"""A data loader that performs mini-batch sampling from link information."""
50
+ def __init__(
51
+ self,
52
+ data: Union[Data, HeteroData],
53
+ link_sampler: Any,
54
+ edge_label_index: Any = None,
55
+ edge_label: Optional[Any] = None,
56
+ edge_label_time: Optional[Any] = None,
57
+ neg_sampling: Optional[Any] = None,
58
+ neg_sampling_ratio: Optional[Union[int, float]] = None,
59
+ transform: Optional[Callable] = None,
60
+ transform_sampler_output: Optional[Callable] = None,
61
+ filter_per_worker: Optional[bool] = None,
62
+ custom_cls: Optional[Any] = None,
63
+ input_id: Optional[Any] = None,
64
+ **kwargs,
65
+ ):
66
+ if filter_per_worker is None:
67
+ filter_per_worker = infer_filter_per_worker(data)
68
+
69
+ kwargs.pop('dataset', None)
70
+ kwargs.pop('collate_fn', None)
71
+
72
+ input_type, edge_label_index = get_edge_label_index(data, edge_label_index)
73
+
74
+ self.data = data
75
+ self.link_sampler = link_sampler
76
+ self.neg_sampling = neg_sampling
77
+ self.neg_sampling_ratio = neg_sampling_ratio
78
+ self.transform = transform
79
+ self.transform_sampler_output = transform_sampler_output
80
+ self.filter_per_worker = filter_per_worker
81
+ self.custom_cls = custom_cls
82
+
83
+ if torch is not None and isinstance(edge_label_index, Tensor):
84
+ row = edge_label_index[0]
85
+ col = edge_label_index[1]
86
+ num_edges = edge_label_index.size(1)
87
+ else:
88
+ np_edges = np.asarray(edge_label_index)
89
+ row = np_edges[0]
90
+ col = np_edges[1]
91
+ num_edges = np_edges.shape[1]
92
+
93
+ self.input_data = EdgeSamplerInput(
94
+ input_id=input_id,
95
+ row=row,
96
+ col=col,
97
+ label=edge_label,
98
+ time=edge_label_time,
99
+ input_type=input_type,
100
+ )
101
+
102
+ iterator = range(num_edges)
103
+
104
+ if torch is not None:
105
+ super().__init__(iterator, collate_fn=self.collate_fn, **kwargs)
106
+ else:
107
+ self.dataset = iterator
108
+ self.collate_fn = self.collate_fn
109
+
110
+ def __call__(self, index: Any) -> Union[Data, HeteroData]:
111
+ out = self.collate_fn(index)
112
+ if not self.filter_per_worker:
113
+ out = self.filter_fn(out)
114
+ return out
115
+
116
+ def collate_fn(self, index: Any) -> Any:
117
+ input_data = self.input_data[index]
118
+ out = self.link_sampler.sample_from_edges(input_data)
119
+ if self.filter_per_worker:
120
+ out = self.filter_fn(out)
121
+ return out
122
+
123
+ def filter_fn(self, out: Any) -> Union[Data, HeteroData]:
124
+ if self.transform_sampler_output:
125
+ out = self.transform_sampler_output(out)
126
+
127
+ if isinstance(out, SamplerOutput):
128
+ perm = getattr(self.link_sampler, 'edge_permutation', None)
129
+ data = filter_data(self.data, out.node, out.row, out.col, out.edge, perm)
130
+
131
+ data.n_id = out.node
132
+ if out.edge is not None:
133
+ data.e_id = out.edge
134
+ data.batch = out.batch
135
+ data.num_sampled_nodes = out.num_sampled_nodes
136
+ data.num_sampled_edges = out.num_sampled_edges
137
+
138
+ meta = out.metadata or (None, None, None)
139
+ data.input_id = meta[0]
140
+ data.edge_label_index = meta[1]
141
+ data.edge_label = meta[2]
142
+
143
+ elif isinstance(out, HeteroSamplerOutput):
144
+ perm = getattr(self.link_sampler, 'edge_permutation', None)
145
+ data = filter_hetero_data(self.data, out.node, out.row, out.col, out.edge, perm)
146
+
147
+ for key, node in out.node.items():
148
+ data[key].n_id = node
149
+
150
+ for key, edge in (out.edge or {}).items():
151
+ if edge is not None:
152
+ data[key].e_id = edge
153
+
154
+ if out.batch is not None:
155
+ data.set_value_dict('batch', out.batch)
156
+ if out.num_sampled_nodes is not None:
157
+ data.set_value_dict('num_sampled_nodes', out.num_sampled_nodes)
158
+ if out.num_sampled_edges is not None:
159
+ data.set_value_dict('num_sampled_edges', out.num_sampled_edges)
160
+
161
+ input_type = self.input_data.input_type
162
+ meta = out.metadata or (None, None, None)
163
+ if input_type is not None:
164
+ data[input_type].input_id = meta[0]
165
+ data[input_type].edge_label_index = meta[1]
166
+ data[input_type].edge_label = meta[2]
167
+ else:
168
+ data = out
169
+
170
+ return data if self.transform is None else self.transform(data)
171
+
172
+ def _get_iterator(self) -> Iterator:
173
+ if self.filter_per_worker:
174
+ return super()._get_iterator()
175
+ return DataLoaderIterator(super()._get_iterator(), self.filter_fn)
176
+
177
+ def __repr__(self) -> str:
178
+ return f'{self.__class__.__name__}()'
179
+