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,167 @@
1
+ from typing import Any, Callable, List, NamedTuple, Optional, Tuple, Union
2
+
3
+ import numpy as np
4
+
5
+ try:
6
+ import torch
7
+ from torch import Tensor
8
+ BaseDataLoader = torch.utils.data.DataLoader
9
+ except ImportError:
10
+ torch = None
11
+ Tensor = type(None)
12
+ BaseDataLoader = object
13
+
14
+ from k3_node.loader.sampler_utils import FastGraph
15
+
16
+
17
+ class EdgeIndex(NamedTuple):
18
+ edge_index: Any
19
+ e_id: Optional[Any]
20
+ size: Tuple[int, int]
21
+
22
+ def to(self, *args, **kwargs):
23
+ if hasattr(self.edge_index, 'to'):
24
+ edge_index = self.edge_index.to(*args, **kwargs)
25
+ e_id = self.e_id.to(*args, **kwargs) if self.e_id is not None and hasattr(self.e_id, 'to') else None
26
+ return EdgeIndex(edge_index, e_id, self.size)
27
+ return self
28
+
29
+
30
+ class Adj(NamedTuple):
31
+ adj_t: Any
32
+ e_id: Optional[Any]
33
+ size: Tuple[int, int]
34
+
35
+ def to(self, *args, **kwargs):
36
+ if hasattr(self.adj_t, 'to'):
37
+ adj_t = self.adj_t.to(*args, **kwargs)
38
+ e_id = self.e_id.to(*args, **kwargs) if self.e_id is not None and hasattr(self.e_id, 'to') else None
39
+ return Adj(adj_t, e_id, self.size)
40
+ return self
41
+
42
+
43
+ class NeighborSampler(BaseDataLoader):
44
+ r"""The legacy layer-by-layer bipartite neighbor sampler from "Inductive Representation
45
+ Learning on Large Graphs".
46
+
47
+ Args:
48
+ edge_index (Tensor): Edge connectivity.
49
+ sizes (List[int]): Number of neighbors to sample per layer.
50
+ node_idx (Tensor, optional): Seed nodes to consider. (default: :obj:`None`)
51
+ num_nodes (int, optional): Number of nodes. (default: :obj:`None`)
52
+ return_e_id (bool, optional): Whether to return edge IDs. (default: :obj:`True`)
53
+ transform (callable, optional): Transform function. (default: :obj:`None`)
54
+ **kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
55
+ """
56
+ def __init__(
57
+ self,
58
+ edge_index: Any,
59
+ sizes: List[int],
60
+ node_idx: Optional[Any] = None,
61
+ num_nodes: Optional[int] = None,
62
+ return_e_id: bool = True,
63
+ transform: Optional[Callable] = None,
64
+ **kwargs,
65
+ ):
66
+ kwargs.pop('dataset', None)
67
+ kwargs.pop('collate_fn', None)
68
+
69
+ self.sizes = sizes
70
+ self.return_e_id = return_e_id
71
+ self.transform = transform
72
+
73
+ is_torch = torch is not None and isinstance(edge_index, Tensor)
74
+ if is_torch:
75
+ np_edge_index = edge_index.detach().cpu().numpy()
76
+ else:
77
+ np_edge_index = np.asarray(edge_index)
78
+
79
+ if num_nodes is None:
80
+ num_nodes = int(np.max(np_edge_index) + 1) if np_edge_index.size > 0 else 0
81
+
82
+ self.num_nodes = num_nodes
83
+ self.graph = FastGraph(np_edge_index, num_nodes=num_nodes)
84
+
85
+ if node_idx is None:
86
+ node_idx = torch.arange(num_nodes) if is_torch else np.arange(num_nodes)
87
+ elif is_torch and isinstance(node_idx, Tensor) and node_idx.dtype == torch.bool:
88
+ node_idx = node_idx.nonzero(as_tuple=False).view(-1)
89
+ elif isinstance(node_idx, np.ndarray) and node_idx.dtype == bool:
90
+ node_idx = np.nonzero(node_idx)[0]
91
+
92
+ self.node_idx = node_idx
93
+ idx_list = node_idx.tolist() if hasattr(node_idx, 'tolist') else list(node_idx)
94
+
95
+ if torch is not None:
96
+ super().__init__(idx_list, collate_fn=self.sample, **kwargs)
97
+ else:
98
+ self.dataset = idx_list
99
+ self.collate_fn = self.sample
100
+
101
+ def sample(self, batch: Any) -> Any:
102
+ is_torch = torch is not None
103
+ if is_torch and not isinstance(batch, Tensor):
104
+ batch = torch.tensor(batch, dtype=torch.long)
105
+ batch_size = batch.numel()
106
+ curr_nodes = batch.tolist()
107
+ else:
108
+ batch_size = len(batch)
109
+ curr_nodes = list(batch)
110
+
111
+ adjs = []
112
+ n_id = list(curr_nodes)
113
+ node_map = {n: i for i, n in enumerate(n_id)}
114
+
115
+ for size in self.sizes:
116
+ rows, cols, e_ids = [], [], []
117
+ targets = list(curr_nodes)
118
+ next_curr_nodes = []
119
+
120
+ for target in targets:
121
+ srcs, edges = self.graph.get_neighbors(target)
122
+ count = len(srcs)
123
+ if count == 0:
124
+ continue
125
+
126
+ if size == -1 or size >= count:
127
+ chosen = np.arange(count)
128
+ else:
129
+ chosen = np.random.choice(count, size=size, replace=False)
130
+
131
+ for s, e in zip(srcs[chosen], edges[chosen]):
132
+ s = int(s)
133
+ e = int(e)
134
+ if s not in node_map:
135
+ node_map[s] = len(n_id)
136
+ n_id.append(s)
137
+ next_curr_nodes.append(s)
138
+
139
+ rows.append(node_map[s])
140
+ cols.append(node_map[target])
141
+ e_ids.append(e)
142
+
143
+ curr_nodes = list(n_id)
144
+
145
+ if len(rows) > 0:
146
+ edge_index = np.stack([np.array(rows, dtype=np.int64), np.array(cols, dtype=np.int64)], axis=0)
147
+ edge_ids = np.array(e_ids, dtype=np.int64) if self.return_e_id else None
148
+ else:
149
+ edge_index = np.empty((2, 0), dtype=np.int64)
150
+ edge_ids = np.empty(0, dtype=np.int64) if self.return_e_id else None
151
+
152
+ bipartite_size = (len(n_id), len(targets))
153
+
154
+ if is_torch:
155
+ edge_index = torch.from_numpy(edge_index)
156
+ if edge_ids is not None:
157
+ edge_ids = torch.from_numpy(edge_ids)
158
+
159
+ adjs.append(EdgeIndex(edge_index, edge_ids, bipartite_size))
160
+
161
+ adjs = adjs[0] if len(adjs) == 1 else adjs[::-1]
162
+ out_n_id = torch.tensor(n_id, dtype=torch.long) if is_torch else np.array(n_id, dtype=np.int64)
163
+ out = (batch_size, out_n_id, adjs)
164
+ return self.transform(*out) if self.transform is not None else out
165
+
166
+ def __repr__(self) -> str:
167
+ return f'{self.__class__.__name__}(sizes={self.sizes})'
@@ -0,0 +1,185 @@
1
+ from dataclasses import dataclass
2
+ from typing import Any, Callable, Dict, Iterator, List, Optional, Tuple, Union
3
+
4
+ try:
5
+ import torch
6
+ from torch import Tensor
7
+ except ImportError:
8
+ torch = None
9
+ Tensor = type(None)
10
+
11
+ from k3_node.data import Data, HeteroData
12
+ from k3_node.loader.base import BaseDataLoader, DataLoaderIterator
13
+ from k3_node.loader.mixin import AffinityMixin, LogMemoryMixin, MultithreadingMixin
14
+ from k3_node.loader.utils import (
15
+ filter_data,
16
+ filter_hetero_data,
17
+ get_input_nodes,
18
+ infer_filter_per_worker,
19
+ )
20
+ from k3_node.loader.keras_dataset import loader_bases
21
+
22
+
23
+ @dataclass
24
+ class NodeSamplerInput:
25
+ input_id: Optional[Any]
26
+ node: Any
27
+ time: Optional[Any] = None
28
+ input_type: Optional[str] = None
29
+
30
+ def __getitem__(self, index: Any) -> 'NodeSamplerInput':
31
+ if torch is not None and not isinstance(index, Tensor):
32
+ index = torch.as_tensor(index, dtype=torch.long)
33
+ return NodeSamplerInput(
34
+ input_id=self.input_id[index] if self.input_id is not None else index,
35
+ node=self.node[index],
36
+ time=self.time[index] if self.time is not None else None,
37
+ input_type=self.input_type,
38
+ )
39
+
40
+
41
+ @dataclass
42
+ class SamplerOutput:
43
+ node: Any
44
+ row: Any
45
+ col: Any
46
+ edge: Optional[Any] = None
47
+ batch: Optional[Any] = None
48
+ num_sampled_nodes: Optional[List[int]] = None
49
+ num_sampled_edges: Optional[List[int]] = None
50
+ orig_row: Optional[Any] = None
51
+ orig_col: Optional[Any] = None
52
+ metadata: Optional[Any] = None
53
+
54
+
55
+ @dataclass
56
+ class HeteroSamplerOutput:
57
+ node: Dict[str, Any]
58
+ row: Dict[Tuple[str, str, str], Any]
59
+ col: Dict[Tuple[str, str, str], Any]
60
+ edge: Dict[Tuple[str, str, str], Optional[Any]]
61
+ batch: Optional[Dict[str, Any]] = None
62
+ num_sampled_nodes: Optional[Dict[str, List[int]]] = None
63
+ num_sampled_edges: Optional[Dict[Tuple[str, str, str], List[int]]] = None
64
+ orig_row: Optional[Dict[Tuple[str, str, str], Any]] = None
65
+ orig_col: Optional[Dict[Tuple[str, str, str], Any]] = None
66
+ metadata: Optional[Any] = None
67
+
68
+
69
+ class NodeLoader(*loader_bases(BaseDataLoader), AffinityMixin, MultithreadingMixin, LogMemoryMixin):
70
+ r"""A data loader that performs mini-batch sampling from node information."""
71
+ def __init__(
72
+ self,
73
+ data: Union[Data, HeteroData],
74
+ node_sampler: Any,
75
+ input_nodes: Any = None,
76
+ input_time: Optional[Any] = None,
77
+ transform: Optional[Callable] = None,
78
+ transform_sampler_output: Optional[Callable] = None,
79
+ filter_per_worker: Optional[bool] = None,
80
+ custom_cls: Optional[Any] = None,
81
+ input_id: Optional[Any] = None,
82
+ **kwargs,
83
+ ):
84
+ if filter_per_worker is None:
85
+ filter_per_worker = infer_filter_per_worker(data)
86
+
87
+ self.data = data
88
+ self.node_sampler = node_sampler
89
+ self.input_nodes = input_nodes
90
+ self.input_time = input_time
91
+ self.transform = transform
92
+ self.transform_sampler_output = transform_sampler_output
93
+ self.filter_per_worker = filter_per_worker
94
+ self.custom_cls = custom_cls
95
+ self.input_id = input_id
96
+
97
+ kwargs.pop('dataset', None)
98
+ kwargs.pop('collate_fn', None)
99
+
100
+ input_type, input_nodes, input_id = get_input_nodes(data, input_nodes, input_id)
101
+
102
+ self.input_data = NodeSamplerInput(
103
+ input_id=input_id,
104
+ node=input_nodes,
105
+ time=input_time,
106
+ input_type=input_type,
107
+ )
108
+
109
+ num_inputs = input_nodes.size(0) if hasattr(input_nodes, 'size') else len(input_nodes)
110
+ iterator = range(num_inputs)
111
+
112
+ if torch is not None:
113
+ super().__init__(iterator, collate_fn=self.collate_fn, **kwargs)
114
+ else:
115
+ self.dataset = iterator
116
+ self.collate_fn = self.collate_fn
117
+
118
+ def __call__(self, index: Any) -> Union[Data, HeteroData]:
119
+ out = self.collate_fn(index)
120
+ if not self.filter_per_worker:
121
+ out = self.filter_fn(out)
122
+ return out
123
+
124
+ def collate_fn(self, index: Any) -> Any:
125
+ input_data = self.input_data[index]
126
+ out = self.node_sampler.sample_from_nodes(input_data)
127
+ if self.filter_per_worker:
128
+ out = self.filter_fn(out)
129
+ return out
130
+
131
+ def filter_fn(self, out: Any) -> Union[Data, HeteroData]:
132
+ if self.transform_sampler_output:
133
+ out = self.transform_sampler_output(out)
134
+
135
+ if isinstance(out, SamplerOutput):
136
+ perm = getattr(self.node_sampler, 'edge_permutation', None)
137
+ data = filter_data(self.data, out.node, out.row, out.col, out.edge, perm)
138
+
139
+ data.n_id = out.node
140
+ if out.edge is not None:
141
+ data.e_id = out.edge
142
+ data.batch = out.batch
143
+ data.num_sampled_nodes = out.num_sampled_nodes
144
+ data.num_sampled_edges = out.num_sampled_edges
145
+
146
+ meta = out.metadata or (out.node, None)
147
+ data.input_id = meta[0]
148
+ data.batch_size = meta[0].size(0) if hasattr(meta[0], 'size') else len(meta[0])
149
+
150
+ elif isinstance(out, HeteroSamplerOutput):
151
+ perm = getattr(self.node_sampler, 'edge_permutation', None)
152
+ data = filter_hetero_data(self.data, out.node, out.row, out.col, out.edge, perm)
153
+
154
+ for key, node in out.node.items():
155
+ data[key].n_id = node
156
+
157
+ for key, edge in (out.edge or {}).items():
158
+ if edge is not None:
159
+ data[key].e_id = edge
160
+
161
+ if out.batch is not None:
162
+ data.set_value_dict('batch', out.batch)
163
+ if out.num_sampled_nodes is not None:
164
+ data.set_value_dict('num_sampled_nodes', out.num_sampled_nodes)
165
+ if out.num_sampled_edges is not None:
166
+ data.set_value_dict('num_sampled_edges', out.num_sampled_edges)
167
+
168
+ input_type = self.input_data.input_type
169
+ meta = out.metadata or (out.node.get(input_type, None), None)
170
+ data[input_type].input_id = meta[0]
171
+ data[input_type].batch_size = meta[0].size(0) if hasattr(meta[0], 'size') else len(meta[0])
172
+
173
+ else:
174
+ data = out
175
+
176
+ return data if self.transform is None else self.transform(data)
177
+
178
+ def _get_iterator(self) -> Iterator:
179
+ if self.filter_per_worker:
180
+ return super()._get_iterator()
181
+ return DataLoaderIterator(super()._get_iterator(), self.filter_fn)
182
+
183
+ def __repr__(self) -> str:
184
+ return f'{self.__class__.__name__}()'
185
+
@@ -0,0 +1,115 @@
1
+ import warnings
2
+ from contextlib import nullcontext
3
+ from functools import partial
4
+ from typing import Any, Optional
5
+
6
+ try:
7
+ import torch
8
+ except ImportError:
9
+ torch = None
10
+ DataLoader = object
11
+
12
+
13
+ class DeviceHelper:
14
+ def __init__(self, device: Optional[Any] = None):
15
+ if torch is None:
16
+ self.device = 'cpu'
17
+ self.is_gpu = False
18
+ self.stream = None
19
+ self.stream_context = nullcontext
20
+ self.module = None
21
+ return
22
+
23
+ with_cuda = torch.cuda.is_available()
24
+ with_xpu = hasattr(torch, 'xpu') and torch.xpu.is_available()
25
+
26
+ if device is None:
27
+ if with_cuda:
28
+ device = 'cuda'
29
+ elif with_xpu:
30
+ device = 'xpu'
31
+ else:
32
+ device = 'cpu'
33
+
34
+ self.device = torch.device(device)
35
+ self.is_gpu = self.device.type in ['cuda', 'xpu']
36
+
37
+ if ((self.device.type == 'cuda' and not with_cuda)
38
+ or (self.device.type == 'xpu' and not with_xpu)):
39
+ warnings.warn(
40
+ f"Requested device '{self.device.type}' is not available, falling back to CPU",
41
+ stacklevel=2,
42
+ )
43
+ self.device = torch.device('cpu')
44
+ self.is_gpu = False
45
+
46
+ self.stream = None
47
+ self.stream_context = nullcontext
48
+ self.module = getattr(torch, self.device.type) if self.is_gpu else None
49
+
50
+ def maybe_init_stream(self) -> None:
51
+ if self.is_gpu and self.module is not None:
52
+ self.stream = self.module.Stream()
53
+ self.stream_context = partial(self.module.stream, stream=self.stream)
54
+
55
+ def maybe_wait_stream(self) -> None:
56
+ if self.stream is not None and self.module is not None:
57
+ self.module.current_stream().wait_stream(self.stream)
58
+
59
+
60
+ class PrefetchLoader:
61
+ r"""A prefetcher class for asynchronously transferring data of a DataLoader
62
+ from host memory to device memory.
63
+
64
+ Args:
65
+ loader (DataLoader): The data loader.
66
+ device (torch.device, optional): The device to load the data to. (default: :obj:`None`)
67
+ """
68
+ def __init__(
69
+ self,
70
+ loader: Any,
71
+ device: Optional[Any] = None,
72
+ ):
73
+ self.loader = loader
74
+ self.device_helper = DeviceHelper(device)
75
+
76
+ def non_blocking_transfer(self, batch: Any) -> Any:
77
+ if not self.device_helper.is_gpu:
78
+ return batch
79
+ if isinstance(batch, (list, tuple)):
80
+ return type(batch)(self.non_blocking_transfer(v) for v in batch)
81
+ if isinstance(batch, dict):
82
+ return {k: self.non_blocking_transfer(v) for k, v in batch.items()}
83
+
84
+ if hasattr(batch, 'pin_memory'):
85
+ batch = batch.pin_memory()
86
+ if hasattr(batch, 'to'):
87
+ return batch.to(self.device_helper.device, non_blocking=True)
88
+ return batch
89
+
90
+ def __iter__(self) -> Any:
91
+ first = True
92
+ self.device_helper.maybe_init_stream()
93
+
94
+ batch = None
95
+ for next_batch in self.loader:
96
+ with self.device_helper.stream_context():
97
+ next_batch = self.non_blocking_transfer(next_batch)
98
+
99
+ if not first:
100
+ yield batch
101
+ else:
102
+ first = False
103
+
104
+ self.device_helper.maybe_wait_stream()
105
+ batch = next_batch
106
+
107
+ if batch is not None:
108
+ yield batch
109
+
110
+ def __len__(self) -> int:
111
+ return len(self.loader)
112
+
113
+ def __repr__(self) -> str:
114
+ return f'{self.__class__.__name__}({self.loader})'
115
+
@@ -0,0 +1,89 @@
1
+ import math
2
+ from typing import Union
3
+
4
+ import numpy as np
5
+
6
+ try:
7
+ import torch
8
+ import torch.utils.data
9
+ from torch import Tensor
10
+ BaseDataLoader = torch.utils.data.DataLoader
11
+ except ImportError:
12
+ torch = None
13
+ Tensor = type(None)
14
+ BaseDataLoader = object
15
+
16
+ from k3_node.data import Data, HeteroData
17
+ from k3_node.loader.keras_dataset import loader_bases
18
+
19
+
20
+ class RandomNodeLoader(*loader_bases(BaseDataLoader)):
21
+ r"""A data loader that randomly samples nodes within a graph and returns
22
+ their induced subgraph.
23
+
24
+ Args:
25
+ data (Data or HeteroData): The graph data object.
26
+ num_parts (int): The number of partitions.
27
+ **kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
28
+ """
29
+ def __init__(
30
+ self,
31
+ data: Union[Data, HeteroData],
32
+ num_parts: int,
33
+ **kwargs,
34
+ ):
35
+ self.data = data
36
+ self.num_parts = num_parts
37
+
38
+ if isinstance(data, HeteroData):
39
+ node_dict = {}
40
+ total = 0
41
+ for node_type in data.node_types:
42
+ count = data[node_type].num_nodes
43
+ node_dict[node_type] = (total, total + count)
44
+ total += count
45
+ self.node_dict = node_dict
46
+ self.num_nodes = total
47
+ else:
48
+ self.edge_index = data.edge_index
49
+ self.num_nodes = data.num_nodes
50
+
51
+ kwargs.pop('dataset', None)
52
+ kwargs.pop('batch_size', None)
53
+ kwargs.pop('collate_fn', None)
54
+
55
+ batch_size = math.ceil(self.num_nodes / num_parts) if num_parts > 0 else self.num_nodes
56
+
57
+ if torch is not None:
58
+ super().__init__(
59
+ range(self.num_nodes),
60
+ batch_size=batch_size,
61
+ collate_fn=self.collate_fn,
62
+ **kwargs,
63
+ )
64
+ else:
65
+ self.dataset = range(self.num_nodes)
66
+ self.batch_size = batch_size
67
+ self.collate_fn = self.collate_fn
68
+
69
+ def collate_fn(self, index):
70
+ if torch is not None and not isinstance(index, Tensor):
71
+ index = torch.tensor(index, dtype=torch.long)
72
+ elif torch is None:
73
+ index = np.asarray(index, dtype=np.int64)
74
+
75
+ if isinstance(self.data, Data):
76
+ return self.data.subgraph(index)
77
+
78
+ elif isinstance(self.data, HeteroData):
79
+ node_dict = {}
80
+ for key, (start, end) in self.node_dict.items():
81
+ if torch is not None and isinstance(index, Tensor):
82
+ mask = (index >= start) & (index < end)
83
+ node_dict[key] = index[mask] - start
84
+ else:
85
+ idx_np = np.asarray(index)
86
+ mask = (idx_np >= start) & (idx_np < end)
87
+ node_dict[key] = idx_np[mask] - start
88
+ return self.data.subgraph(node_dict)
89
+