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,107 @@
1
+ from typing import Callable, Optional, Union, Tuple
2
+ import keras
3
+ from keras import layers, ops
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+ from k3_node.layers.pool.knn import knn_graph
7
+
8
+
9
+ class EdgeConv(MessagePassing):
10
+ r"""The edge convolutional operator from the `"Dynamic Graph CNN for
11
+ Learning on Point Clouds" <https://arxiv.org/abs/1801.07829>`_ paper.
12
+
13
+ Args:
14
+ nn: A neural network :math:`h_{\mathbf{\Theta}}` that maps
15
+ pair-wise node features to new edge representations.
16
+ aggr: The aggregation scheme to use (``"max"``, ``"mean"``, ``"sum"``).
17
+ (default: ``"max"``)
18
+
19
+ Example:
20
+ ```python
21
+ import numpy as np
22
+ import keras
23
+ from k3_node.layers import EdgeConv
24
+
25
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
26
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
27
+
28
+ nn = keras.Sequential([keras.layers.Dense(16, activation="relu"), keras.layers.Dense(16)])
29
+ layer = EdgeConv(nn)
30
+ out = layer(x, edge_index)
31
+ print(tuple(out.shape)) # (10, 16)
32
+ ```
33
+ """
34
+
35
+ def __init__(self, nn: Callable, aggr: str = "max", **kwargs):
36
+ super().__init__(aggr=aggr, **kwargs)
37
+ self.nn = nn
38
+
39
+ def build(self, input_shape):
40
+ if hasattr(self.nn, "build") and not getattr(self.nn, "built", False):
41
+ if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
42
+ c = input_shape[0][-1]
43
+ elif isinstance(input_shape, (tuple, list)):
44
+ c = input_shape[-1]
45
+ else:
46
+ c = None
47
+ nn_shape = (None, 2 * c) if c is not None else None
48
+ self.nn.build(nn_shape)
49
+ self.built = True
50
+
51
+ def call(self, x, edge_index=None, training=None, **kwargs):
52
+ if edge_index is None and isinstance(x, (tuple, list)):
53
+ x, edge_index = x[0], x[1]
54
+
55
+ if not isinstance(x, (tuple, list)):
56
+ x_src, x_dst = x, x
57
+ else:
58
+ x_src, x_dst = x[0], x[1]
59
+
60
+ return self.propagate(edge_index, x=(x_src, x_dst), training=training, **kwargs)
61
+
62
+ def message(self, x_i, x_j, training=None):
63
+ h = ops.concatenate([x_i, x_j - x_i], axis=-1)
64
+ # Forward `training` explicitly: Keras does not propagate it to nested layers on JAX.
65
+ return self.nn(h, training=training) if isinstance(self.nn, keras.layers.Layer) else self.nn(h)
66
+
67
+
68
+ class DynamicEdgeConv(EdgeConv):
69
+ r"""The dynamic edge convolutional operator from the `"Dynamic Graph CNN
70
+ for Learning on Point Clouds" <https://arxiv.org/abs/1801.07829>`_ paper,
71
+ which dynamically constructs a graph using :math:`k`-NN at each layer.
72
+
73
+ Args:
74
+ nn: A neural network :math:`h_{\mathbf{\Theta}}`.
75
+ k: Number of nearest neighbors. (default: ``6``)
76
+ aggr: The aggregation scheme to use (``"max"``, ``"mean"``, ``"sum"``).
77
+ (default: ``"max"``)
78
+ num_workers: Number of workers (ignored in Keras backend).
79
+
80
+ Example:
81
+ ```python
82
+ import numpy as np
83
+ import keras
84
+ from k3_node.layers import DynamicEdgeConv
85
+
86
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
87
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
88
+
89
+ nn = keras.Sequential([keras.layers.Dense(16, activation="relu"), keras.layers.Dense(16)])
90
+ layer = DynamicEdgeConv(nn, k=3)
91
+ out = layer(x) # k-NN graph is built from x
92
+ print(tuple(out.shape)) # (10, 16)
93
+ ```
94
+ """
95
+
96
+ def __init__(self, nn: Callable, k: int = 6, aggr: str = "max", num_workers: int = 1, **kwargs):
97
+ super().__init__(nn=nn, aggr=aggr, **kwargs)
98
+ self.k = k
99
+
100
+ def call(self, x, batch=None, training=None, **kwargs):
101
+ if isinstance(x, (tuple, list)):
102
+ x_src = x[0]
103
+ else:
104
+ x_src = x
105
+
106
+ edge_index = knn_graph(x_src, k=self.k, batch=batch, loop=False, flow=self.flow)
107
+ return super().call(x, edge_index=edge_index, training=training, **kwargs)
@@ -0,0 +1,155 @@
1
+ from typing import Optional, List
2
+ from keras import ops
3
+ from keras.layers import Dense
4
+
5
+ from k3_node.layers.conv.message_passing import MessagePassing
6
+ from k3_node.layers.conv.utils import gcn_norm, scatter
7
+
8
+
9
+ class EGConv(MessagePassing):
10
+ r"""The Efficient Graph Convolution from the `"Adaptive Filters and
11
+ Aggregator Fusion for Efficient Graph Convolutions"
12
+ <https://arxiv.org/abs/2104.01481>`_ paper.
13
+
14
+ Example:
15
+ ```python
16
+ import numpy as np
17
+ from k3_node.layers import EGConv
18
+
19
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
20
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
21
+
22
+ layer = EGConv(in_channels=8, out_channels=16, num_heads=2)
23
+ out = layer(x, edge_index)
24
+ print(tuple(out.shape)) # (10, 16)
25
+ ```
26
+ """
27
+ def __init__(
28
+ self,
29
+ in_channels: int,
30
+ out_channels: int,
31
+ aggregators: Optional[List[str]] = None,
32
+ num_heads: int = 8,
33
+ num_bases: int = 4,
34
+ cached: bool = False,
35
+ add_self_loops: bool = True,
36
+ bias: bool = True,
37
+ **kwargs,
38
+ ):
39
+ super().__init__(node_dim=0, **kwargs)
40
+
41
+ if out_channels % num_heads != 0:
42
+ raise ValueError(
43
+ f"'out_channels' ({out_channels}) must be divisible by num_heads ({num_heads})"
44
+ )
45
+
46
+ self.in_channels = in_channels
47
+ self.out_channels = out_channels
48
+ self.num_heads = num_heads
49
+ self.num_bases = num_bases
50
+ self.cached = cached
51
+ self.add_self_loops = add_self_loops
52
+ self.aggregators = aggregators or ["symnorm"]
53
+ self.use_bias = bias
54
+
55
+ self.bases_lin = Dense(
56
+ (out_channels // num_heads) * num_bases, use_bias=False
57
+ )
58
+ self.comb_lin = Dense(
59
+ num_heads * num_bases * len(self.aggregators), use_bias=True
60
+ )
61
+
62
+ if bias:
63
+ self.bias = self.add_weight(
64
+ shape=(out_channels,),
65
+ initializer="zeros",
66
+ name="bias",
67
+ )
68
+ else:
69
+ self.bias = None
70
+
71
+ def build(self, input_shape=None):
72
+ self.bases_lin.build((None, self.in_channels))
73
+ self.comb_lin.build((None, self.in_channels))
74
+ self.built = True
75
+
76
+ def call(self, inputs, edge_index=None, **kwargs):
77
+ if edge_index is None:
78
+ if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
79
+ x, edge_index = inputs
80
+ else:
81
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
82
+ else:
83
+ x = inputs
84
+
85
+ if not self.built:
86
+ self.build()
87
+
88
+ num_nodes = ops.shape(x)[0]
89
+ symnorm_weight = None
90
+ if "symnorm" in self.aggregators:
91
+ edge_index, symnorm_weight = gcn_norm(
92
+ edge_index,
93
+ edge_weight=None,
94
+ num_nodes=num_nodes,
95
+ add_self_loops=self.add_self_loops,
96
+ dtype=x.dtype,
97
+ )
98
+
99
+ bases = self.bases_lin(x)
100
+ weightings = self.comb_lin(x)
101
+
102
+ aggregated = self.propagate(
103
+ edge_index,
104
+ x=bases,
105
+ symnorm_weight=symnorm_weight,
106
+ size=(num_nodes, num_nodes),
107
+ )
108
+
109
+ weightings = ops.reshape(
110
+ weightings,
111
+ (-1, self.num_heads, self.num_bases * len(self.aggregators)),
112
+ )
113
+ aggregated = ops.reshape(
114
+ aggregated,
115
+ (
116
+ -1,
117
+ len(self.aggregators) * self.num_bases,
118
+ self.out_channels // self.num_heads,
119
+ ),
120
+ )
121
+
122
+ out = ops.matmul(weightings, aggregated)
123
+ out = ops.reshape(out, (-1, self.out_channels))
124
+
125
+ if self.bias is not None:
126
+ out = out + self.bias
127
+
128
+ return out
129
+
130
+ def message(self, x_j):
131
+ return x_j
132
+
133
+ def aggregate(self, inputs, edge_index=None, index=None, dim_size=None, symnorm_weight=None, **kwargs):
134
+ if index is None and edge_index is not None:
135
+ index = edge_index[1]
136
+
137
+ outs = []
138
+ for aggr in self.aggregators:
139
+ if aggr == "symnorm":
140
+ inp = inputs if symnorm_weight is None else inputs * ops.expand_dims(symnorm_weight, -1)
141
+ out = scatter(inp, index, dim=0, dim_size=dim_size, reduce="sum")
142
+ elif aggr in ("var", "std"):
143
+ mean = scatter(inputs, index, dim=0, dim_size=dim_size, reduce="mean")
144
+ mean_sq = scatter(inputs * inputs, index, dim=0, dim_size=dim_size, reduce="mean")
145
+ out = mean_sq - mean * mean
146
+ if aggr == "std":
147
+ out = ops.sqrt(ops.maximum(out, 1e-5))
148
+ else:
149
+ out = scatter(inputs, index, dim=0, dim_size=dim_size, reduce=aggr)
150
+ outs.append(out)
151
+
152
+ if len(outs) > 1:
153
+ return ops.stack(outs, axis=1)
154
+ return outs[0]
155
+
@@ -0,0 +1,107 @@
1
+ from keras import ops
2
+ from keras.layers import Dense
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+ from k3_node.layers.conv.utils import gcn_norm
6
+
7
+
8
+ class FAConv(MessagePassing):
9
+ r"""The Frequency Adaptive Graph Convolution operator from the
10
+ `"Beyond Low-Frequency Information in Graph Convolutional Networks"
11
+ <https://arxiv.org/abs/2101.00797>`_ paper.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ from k3_node.layers import FAConv
17
+
18
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
19
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
20
+
21
+ x_0 = x # initial node representations
22
+ layer = FAConv(channels=8, eps=0.1)
23
+ out = layer(x, x_0, edge_index)
24
+ print(tuple(out.shape)) # (10, 8)
25
+ ```
26
+ """
27
+ def __init__(
28
+ self,
29
+ channels: int,
30
+ eps: float = 0.1,
31
+ dropout: float = 0.0,
32
+ cached: bool = False,
33
+ add_self_loops: bool = True,
34
+ normalize: bool = True,
35
+ **kwargs,
36
+ ):
37
+ kwargs.setdefault("aggr", "add")
38
+ super().__init__(**kwargs)
39
+
40
+ self.channels = channels
41
+ self.eps = eps
42
+ self.dropout_rate = dropout
43
+ self.cached = cached
44
+ self.add_self_loops = add_self_loops
45
+ self.normalize = normalize
46
+
47
+ self.att_l = Dense(1, use_bias=False)
48
+ self.att_r = Dense(1, use_bias=False)
49
+
50
+ def build(self, input_shape=None):
51
+ self.att_l.build((None, self.channels))
52
+ self.att_r.build((None, self.channels))
53
+ self.built = True
54
+
55
+ def call(self, inputs, x_0=None, edge_index=None, edge_weight=None, **kwargs):
56
+ if edge_index is None:
57
+ if isinstance(inputs, (list, tuple)):
58
+ if len(inputs) == 4:
59
+ x, x_0, edge_index, edge_weight = inputs
60
+ elif len(inputs) == 3:
61
+ x, x_0, edge_index = inputs
62
+ elif len(inputs) == 2:
63
+ x, edge_index = inputs
64
+ x_0 = x
65
+ else:
66
+ raise ValueError(f"Unexpected input length {len(inputs)}")
67
+ else:
68
+ raise ValueError("Expected inputs with edge_index")
69
+ else:
70
+ x = inputs
71
+ if x_0 is None:
72
+ x_0 = x
73
+
74
+ if not self.built:
75
+ self.build()
76
+
77
+ num_nodes = ops.shape(x)[0]
78
+ if self.normalize:
79
+ edge_index, edge_weight = gcn_norm(
80
+ edge_index,
81
+ edge_weight,
82
+ num_nodes=num_nodes,
83
+ add_self_loops=self.add_self_loops,
84
+ dtype=x.dtype,
85
+ )
86
+
87
+ alpha_l = self.att_l(x)
88
+ alpha_r = self.att_r(x)
89
+
90
+ out = self.propagate(
91
+ edge_index,
92
+ x=x,
93
+ alpha=(alpha_l, alpha_r),
94
+ edge_weight=edge_weight,
95
+ )
96
+
97
+ if self.eps != 0.0:
98
+ out = out + self.eps * x_0
99
+
100
+ return out
101
+
102
+ def message(self, x_j, alpha_j, alpha_i, edge_weight=None):
103
+ alpha = ops.tanh(alpha_j + alpha_i)
104
+ if edge_weight is not None:
105
+ alpha = alpha * ops.expand_dims(edge_weight, -1)
106
+ return x_j * alpha
107
+
@@ -0,0 +1,126 @@
1
+ from keras import ops
2
+ from keras.layers import Dense
3
+
4
+ from k3_node.layers.conv.message_passing import MessagePassing
5
+ from k3_node.layers.conv.utils import (
6
+ add_self_loops,
7
+ degree,
8
+ extend_mask_for_self_loops,
9
+ remove_self_loops_masked,
10
+ )
11
+ from k3_node.ops.segment import segment_sum
12
+
13
+
14
+ class FeaStConv(MessagePassing):
15
+ r"""The (fault-tolerant) feature-steered graph convolution operator from
16
+ the `"FeaStNet: Feature-Steered Graph Convolutions for 3D Shape Analysis"
17
+ <https://arxiv.org/abs/1706.05206>`_ paper.
18
+
19
+ Example:
20
+ ```python
21
+ import numpy as np
22
+ from k3_node.layers import FeaStConv
23
+
24
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
25
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
26
+
27
+ layer = FeaStConv(in_channels=8, out_channels=16, heads=2)
28
+ out = layer(x, edge_index)
29
+ print(tuple(out.shape)) # (10, 16)
30
+ ```
31
+ """
32
+ def __init__(
33
+ self,
34
+ in_channels: int,
35
+ out_channels: int,
36
+ heads: int = 1,
37
+ add_self_loops: bool = True,
38
+ bias: bool = True,
39
+ **kwargs,
40
+ ):
41
+ kwargs.setdefault("aggr", "mean")
42
+ super().__init__(node_dim=0, **kwargs)
43
+
44
+ self.in_channels = in_channels
45
+ self.out_channels = out_channels
46
+ self.heads = heads
47
+ self.add_self_loops = add_self_loops
48
+ self.use_bias = bias
49
+
50
+ self.lin = Dense(heads * out_channels, use_bias=False)
51
+ self.u = self.add_weight(
52
+ shape=(in_channels, heads),
53
+ initializer="glorot_uniform",
54
+ name="u",
55
+ )
56
+ self.c = self.add_weight(
57
+ shape=(heads,),
58
+ initializer="zeros",
59
+ name="c",
60
+ )
61
+
62
+ if bias:
63
+ self.bias = self.add_weight(
64
+ shape=(out_channels,),
65
+ initializer="zeros",
66
+ name="bias",
67
+ )
68
+ else:
69
+ self.bias = None
70
+
71
+ def build(self, input_shape=None):
72
+ self.lin.build((None, self.in_channels))
73
+ self.built = True
74
+
75
+ def call(self, inputs, edge_index=None, **kwargs):
76
+ if edge_index is None:
77
+ if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
78
+ x, edge_index = inputs
79
+ else:
80
+ raise ValueError("Expected (x, edge_index) or x and edge_index")
81
+ else:
82
+ x = inputs
83
+
84
+ if not self.built:
85
+ self.build()
86
+
87
+ if isinstance(x, (list, tuple)):
88
+ x_src, x_dst = x
89
+ else:
90
+ x_src = x_dst = x
91
+
92
+ num_nodes = ops.shape(x_dst)[0]
93
+ keep_mask = None
94
+ if self.add_self_loops:
95
+ edge_index, _, keep_mask = remove_self_loops_masked(edge_index)
96
+ edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
97
+ keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
98
+
99
+ out = self.propagate(
100
+ edge_index,
101
+ x=(x_src, x_dst),
102
+ keep_mask=keep_mask,
103
+ size=(ops.shape(x_src)[0], num_nodes),
104
+ )
105
+ if keep_mask is not None and self.aggr == "mean":
106
+ # Masked messages are zero but still counted by the mean; rescale to the kept count.
107
+ col = ops.cast(edge_index[1], "int32")
108
+ count_all = degree(col, num_nodes=num_nodes, dtype=out.dtype)
109
+ count_kept = segment_sum(ops.cast(keep_mask, out.dtype), col, num_segments=num_nodes)
110
+ out = out * ops.expand_dims(count_all / ops.maximum(count_kept, 1.0), -1)
111
+
112
+ if self.bias is not None:
113
+ out = out + self.bias
114
+
115
+ return out
116
+
117
+ def message(self, x_i, x_j, keep_mask=None):
118
+ q = ops.matmul(x_j - x_i, self.u) + self.c
119
+ q = ops.softmax(q, axis=-1)
120
+ x_j_mapped = ops.reshape(self.lin(x_j), (-1, self.heads, self.out_channels))
121
+ msg = ops.sum(x_j_mapped * ops.expand_dims(q, -1), axis=1)
122
+ # For max/min a duplicated self-loop message is harmless, so only sum/mean need masking.
123
+ if keep_mask is not None and self.aggr in ("add", "sum", "mean"):
124
+ msg = msg * ops.expand_dims(ops.cast(keep_mask, msg.dtype), -1)
125
+ return msg
126
+
@@ -0,0 +1,143 @@
1
+ import copy
2
+ from typing import Optional, Union, Tuple, Callable
3
+ from keras import ops, activations
4
+ from keras.layers import Dense
5
+
6
+ from k3_node.layers.conv.message_passing import MessagePassing
7
+
8
+
9
+ class FiLMConv(MessagePassing):
10
+ r"""The FiLM graph convolutional operator from the
11
+ `"GNN-FiLM: Graph Neural Networks with Feature-wise Linear Modulation"
12
+ <https://arxiv.org/abs/1906.12192>`_ paper.
13
+
14
+ Example:
15
+ ```python
16
+ import numpy as np
17
+ from k3_node.layers import FiLMConv
18
+
19
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
20
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
21
+
22
+ layer = FiLMConv(in_channels=8, out_channels=16)
23
+ out = layer(x, edge_index)
24
+ print(tuple(out.shape)) # (10, 16)
25
+ ```
26
+ """
27
+ def __init__(
28
+ self,
29
+ in_channels: Union[int, Tuple[int, int]],
30
+ out_channels: int,
31
+ num_relations: int = 1,
32
+ nn: Optional[Callable] = None,
33
+ act: Optional[Union[str, Callable]] = "relu",
34
+ aggr: str = "mean",
35
+ **kwargs,
36
+ ):
37
+ super().__init__(aggr=aggr, **kwargs)
38
+
39
+ self.in_channels = in_channels
40
+ self.out_channels = out_channels
41
+ self.num_relations = max(num_relations, 1)
42
+ self.act = activations.get(act) if act is not None else None
43
+
44
+ if isinstance(in_channels, int):
45
+ self.in_channels_l = in_channels
46
+ self.in_channels_r = in_channels
47
+ else:
48
+ self.in_channels_l, self.in_channels_r = in_channels
49
+
50
+ self.lins = [
51
+ Dense(out_channels, use_bias=False) for _ in range(self.num_relations)
52
+ ]
53
+ self.films = []
54
+ for _ in range(self.num_relations):
55
+ if nn is None:
56
+ self.films.append(Dense(2 * out_channels, use_bias=True))
57
+ else:
58
+ self.films.append(copy.deepcopy(nn))
59
+
60
+ self.lin_skip = Dense(out_channels, use_bias=False)
61
+ if nn is None:
62
+ self.film_skip = Dense(2 * out_channels, use_bias=False)
63
+ else:
64
+ self.film_skip = copy.deepcopy(nn)
65
+
66
+ def build(self, input_shape=None):
67
+ for lin in self.lins:
68
+ lin.build((None, self.in_channels_l))
69
+ for film in self.films:
70
+ film.build((None, self.in_channels_r))
71
+ self.lin_skip.build((None, self.in_channels_r))
72
+ self.film_skip.build((None, self.in_channels_r))
73
+ self.built = True
74
+
75
+ def call(self, inputs, edge_index=None, edge_type=None, **kwargs):
76
+ if edge_index is None:
77
+ if isinstance(inputs, (list, tuple)):
78
+ if len(inputs) == 3:
79
+ x, edge_index, edge_type = inputs
80
+ elif len(inputs) == 2:
81
+ x, edge_index = inputs
82
+ else:
83
+ raise ValueError(f"Unexpected input length {len(inputs)}")
84
+ else:
85
+ raise ValueError("Expected inputs with edge_index")
86
+ else:
87
+ x = inputs
88
+
89
+ if not self.built:
90
+ self.build()
91
+
92
+ if isinstance(x, (list, tuple)):
93
+ x_l, x_r = x
94
+ else:
95
+ x_l = x_r = x
96
+
97
+ edge_index = ops.cast(edge_index, "int32")
98
+ if edge_type is not None:
99
+ edge_type = ops.cast(edge_type, "int32")
100
+
101
+ # Skip connection
102
+ film_s = self.film_skip(x_r)
103
+ beta_s, gamma_s = ops.split(film_s, 2, axis=-1)
104
+ out = gamma_s * self.lin_skip(x_r) + beta_s
105
+ if self.act is not None:
106
+ out = self.act(out)
107
+
108
+ # Message passing per relation
109
+ num_nodes = ops.shape(x_r)[0]
110
+ size = (ops.shape(x_l)[0], num_nodes)
111
+
112
+ for i in range(self.num_relations):
113
+ if edge_type is not None:
114
+ mask = ops.equal(edge_type, i)
115
+ where_mask = ops.where(mask)
116
+ idx = where_mask[0] if isinstance(where_mask, (list, tuple)) else where_mask
117
+ idx = ops.reshape(idx, (-1,))
118
+ idx = ops.cast(idx, "int32")
119
+ if ops.shape(idx)[0] == 0:
120
+ continue
121
+ edge_index_i = ops.take(edge_index, idx, axis=1)
122
+ else:
123
+ edge_index_i = edge_index
124
+
125
+ film_val = self.films[i](x_r)
126
+ beta, gamma = ops.split(film_val, 2, axis=-1)
127
+ h = self.propagate(
128
+ edge_index_i,
129
+ x=self.lins[i](x_l),
130
+ beta=beta,
131
+ gamma=gamma,
132
+ size=size,
133
+ )
134
+ out = out + h
135
+
136
+ return out
137
+
138
+ def message(self, x_j, beta_i, gamma_i):
139
+ out = gamma_i * x_j + beta_i
140
+ if self.act is not None:
141
+ out = self.act(out)
142
+ return out
143
+