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,234 @@
1
+ from typing import Dict, List, Optional, Tuple
2
+ import numpy as np
3
+ import keras
4
+ from keras import ops
5
+
6
+ EdgeType = Tuple[str, str, str]
7
+ NodeType = str
8
+
9
+
10
+ class MetaPath2Vec(keras.layers.Layer):
11
+ r"""The MetaPath2Vec model from the `"metapath2vec: Scalable Representation
12
+ Learning for Heterogeneous Networks"
13
+ <https://ericdongyx.github.io/papers/
14
+ KDD17-dong-chawla-swami-metapath2vec.pdf>`_ paper where random walks based
15
+ on a given :obj:`metapath` are sampled in a heterogeneous graph, and node
16
+ embeddings are learned via negative sampling optimization.
17
+
18
+ Args:
19
+ edge_index_dict (Dict[Tuple[str, str, str], Tensor]): Dictionary
20
+ holding edge indices for each edge type.
21
+ embedding_dim (int): The size of each embedding vector.
22
+ metapath (List[Tuple[str, str, str]]): The sequence of edge types
23
+ denoting the metapath.
24
+ walk_length (int): The walk length.
25
+ context_size (int): The context size considered for positive samples.
26
+ walks_per_node (int, optional): The number of walks to sample for each node.
27
+ (default: :obj:`1`)
28
+ num_negative_samples (int, optional): The number of negative samples.
29
+ (default: :obj:`1`)
30
+ num_nodes_dict (Dict[str, int], optional): The number of nodes for each
31
+ node type. (default: :obj:`None`)
32
+
33
+ Example:
34
+ ```python
35
+ import numpy as np
36
+ from k3_node.models import MetaPath2Vec
37
+
38
+ edge_index_dict = {
39
+ ("author", "writes", "paper"): np.array([[0, 1, 1], [0, 0, 1]]),
40
+ ("paper", "written_by", "author"): np.array([[0, 0, 1], [0, 1, 1]]),
41
+ }
42
+ metapath = [("author", "writes", "paper"), ("paper", "written_by", "author")]
43
+ model = MetaPath2Vec(edge_index_dict, embedding_dim=16, metapath=metapath,
44
+ walk_length=2, context_size=2, walks_per_node=2)
45
+ print(tuple(model("author").shape)) # (2, 16): embeddings of all authors
46
+ print(tuple(model("paper").shape)) # (2, 16)
47
+ ```
48
+ """
49
+ def __init__(
50
+ self,
51
+ edge_index_dict: Dict[EdgeType, any],
52
+ embedding_dim: int,
53
+ metapath: List[EdgeType],
54
+ walk_length: int,
55
+ context_size: int,
56
+ walks_per_node: int = 1,
57
+ num_negative_samples: int = 1,
58
+ num_nodes_dict: Optional[Dict[NodeType, int]] = None,
59
+ **kwargs,
60
+ ):
61
+ super().__init__(**kwargs)
62
+
63
+ if num_nodes_dict is None:
64
+ num_nodes_dict = {}
65
+ for keys, edge_index in edge_index_dict.items():
66
+ e_np = ops.convert_to_numpy(edge_index)
67
+ key = keys[0]
68
+ N = int(e_np[0].max() + 1) if e_np.size > 0 else 0
69
+ num_nodes_dict[key] = max(N, num_nodes_dict.get(key, N))
70
+
71
+ key = keys[-1]
72
+ N = int(e_np[1].max() + 1) if e_np.size > 0 else 0
73
+ num_nodes_dict[key] = max(N, num_nodes_dict.get(key, N))
74
+
75
+ # Build adjacency dictionaries
76
+ self.adj_dict = {}
77
+ for keys, edge_index in edge_index_dict.items():
78
+ src_type, _, dst_type = keys
79
+ num_src = num_nodes_dict[src_type]
80
+ adj = [[] for _ in range(num_src)]
81
+ e_np = ops.convert_to_numpy(edge_index).astype(np.int64)
82
+ if e_np.size > 0:
83
+ for s, d in zip(e_np[0], e_np[1]):
84
+ adj[s].append(int(d))
85
+ self.adj_dict[keys] = adj
86
+
87
+ for edge_type1, edge_type2 in zip(metapath[:-1], metapath[1:]):
88
+ if edge_type1[-1] != edge_type2[0]:
89
+ raise ValueError(
90
+ "Found invalid metapath. Ensure that the destination node "
91
+ "type matches with the source node type across all "
92
+ "consecutive edge types."
93
+ )
94
+
95
+ assert walk_length + 1 >= context_size
96
+
97
+ self.embedding_dim = embedding_dim
98
+ self.metapath = metapath
99
+ self.walk_length = walk_length
100
+ self.context_size = context_size
101
+ self.walks_per_node = walks_per_node
102
+ self.num_negative_samples = num_negative_samples
103
+ self.num_nodes_dict = num_nodes_dict
104
+ self.EPS = 1e-15
105
+
106
+ types = sorted(list({x[0] for x in metapath} | {x[-1] for x in metapath}))
107
+
108
+ count = 0
109
+ self.start, self.end = {}, {}
110
+ for key in types:
111
+ self.start[key] = count
112
+ count += num_nodes_dict[key]
113
+ self.end[key] = count
114
+
115
+ offset = [self.start[metapath[0][0]]]
116
+ offset += [self.start[keys[-1]] for keys in metapath] * int(
117
+ (walk_length / len(metapath)) + 1
118
+ )
119
+ offset = offset[: walk_length + 1]
120
+ self.offset = np.array(offset, dtype=np.int64)
121
+
122
+ self.dummy_idx = count
123
+ self.embedding = keras.layers.Embedding(count + 1, embedding_dim)
124
+
125
+ def reset_parameters(self):
126
+ if self.embedding.built:
127
+ self.embedding.embeddings.assign(
128
+ keras.initializers.GlorotUniform()(self.embedding.embeddings.shape)
129
+ )
130
+
131
+ def __call__(self, node_type: str, batch: Optional[any] = None, **kwargs):
132
+ return self.call(node_type, batch)
133
+
134
+ def forward(self, node_type: str, batch: Optional[any] = None, **kwargs):
135
+ return self.call(node_type, batch)
136
+
137
+ def call(self, node_type: str, batch: Optional[any] = None):
138
+ r"""Returns the embeddings for the nodes in :obj:`batch` of type
139
+ :obj:`node_type`.
140
+ """
141
+ start = self.start[node_type]
142
+ end = self.end[node_type]
143
+ if batch is None:
144
+ batch = ops.arange(start, end, dtype="int64")
145
+ else:
146
+ batch = ops.cast(batch, "int64") + start
147
+ return self.embedding(batch)
148
+
149
+ def _pos_sample(self, batch):
150
+ batch_np = ops.convert_to_numpy(batch).astype(np.int64)
151
+ repeated = np.repeat(batch_np, self.walks_per_node)
152
+
153
+ all_walks = []
154
+ for node in repeated:
155
+ walk = [int(node)]
156
+ for i in range(self.walk_length):
157
+ edge_type = self.metapath[i % len(self.metapath)]
158
+ cur = walk[-1]
159
+ nbrs = (
160
+ self.adj_dict[edge_type][cur]
161
+ if cur < len(self.adj_dict[edge_type])
162
+ else []
163
+ )
164
+ if len(nbrs) > 0:
165
+ walk.append(nbrs[np.random.randint(len(nbrs))])
166
+ else:
167
+ walk.append(self.dummy_idx)
168
+ all_walks.append(walk)
169
+
170
+ rw = np.array(all_walks, dtype=np.int64)
171
+ rw = rw + self.offset[None, :]
172
+ rw[rw > self.dummy_idx] = self.dummy_idx
173
+
174
+ walks = []
175
+ num_walks_per_rw = 1 + self.walk_length + 1 - self.context_size
176
+ for j in range(num_walks_per_rw):
177
+ walks.append(rw[:, j : j + self.context_size])
178
+ out = np.concatenate(walks, axis=0) if len(walks) > 0 else rw
179
+ return ops.convert_to_tensor(out, dtype="int64")
180
+
181
+ def _neg_sample(self, batch):
182
+ batch_np = ops.convert_to_numpy(batch).astype(np.int64)
183
+ repeated = np.repeat(
184
+ batch_np, self.walks_per_node * self.num_negative_samples
185
+ )
186
+
187
+ rws = [repeated]
188
+ for i in range(self.walk_length):
189
+ keys = self.metapath[i % len(self.metapath)]
190
+ num_nodes = self.num_nodes_dict[keys[-1]]
191
+ rand = np.random.randint(0, max(num_nodes, 1), size=len(repeated))
192
+ rws.append(rand)
193
+
194
+ rw = np.stack(rws, axis=-1)
195
+ rw = rw + self.offset[None, :]
196
+
197
+ walks = []
198
+ num_walks_per_rw = 1 + self.walk_length + 1 - self.context_size
199
+ for j in range(num_walks_per_rw):
200
+ walks.append(rw[:, j : j + self.context_size])
201
+ out = np.concatenate(walks, axis=0) if len(walks) > 0 else rw
202
+ return ops.convert_to_tensor(out, dtype="int64")
203
+
204
+ def loss(self, pos_rw, neg_rw):
205
+ r"""Computes the loss given positive and negative random walks."""
206
+ # Positive loss
207
+ start = pos_rw[:, 0]
208
+ rest = pos_rw[:, 1:]
209
+ pos_b = ops.shape(pos_rw)[0]
210
+
211
+ h_start = ops.reshape(self.embedding(start), (pos_b, 1, self.embedding_dim))
212
+ h_rest = ops.reshape(
213
+ self.embedding(ops.reshape(rest, (-1,))),
214
+ (pos_b, -1, self.embedding_dim),
215
+ )
216
+
217
+ out = ops.reshape(ops.sum(h_start * h_rest, axis=-1), (-1,))
218
+ pos_loss = -ops.mean(ops.log(ops.sigmoid(out) + self.EPS))
219
+
220
+ # Negative loss
221
+ start = neg_rw[:, 0]
222
+ rest = neg_rw[:, 1:]
223
+ neg_b = ops.shape(neg_rw)[0]
224
+
225
+ h_start = ops.reshape(self.embedding(start), (neg_b, 1, self.embedding_dim))
226
+ h_rest = ops.reshape(
227
+ self.embedding(ops.reshape(rest, (-1,))),
228
+ (neg_b, -1, self.embedding_dim),
229
+ )
230
+
231
+ out = ops.reshape(ops.sum(h_start * h_rest, axis=-1), (-1,))
232
+ neg_loss = -ops.mean(ops.log(1.0 - ops.sigmoid(out) + self.EPS))
233
+
234
+ return pos_loss + neg_loss
k3_node/models/mlp.py ADDED
@@ -0,0 +1,264 @@
1
+ import re
2
+ import inspect
3
+ from typing import List, Optional, Union
4
+
5
+ import keras
6
+ from keras import ops
7
+
8
+ import k3_node.layers.norm as norm_module
9
+
10
+
11
+ def _normalize_string(s: str) -> str:
12
+ return re.sub(r"[_\-\s]", "", s).lower()
13
+
14
+
15
+ def _normalization_resolver(query, *args, **kwargs):
16
+ if query is None:
17
+ return None
18
+ if not isinstance(query, str):
19
+ return query
20
+
21
+ norms = {
22
+ _normalize_string(name): cls
23
+ for name, cls in vars(norm_module).items()
24
+ if isinstance(cls, type)
25
+ }
26
+ key = _normalize_string(query)
27
+ if key in norms:
28
+ return norms[key](*args, **kwargs)
29
+ if key + "norm" in norms:
30
+ return norms[key + "norm"](*args, **kwargs)
31
+ if key.endswith("norm") and key[:-4] in norms:
32
+ return norms[key[:-4]](*args, **kwargs)
33
+ raise ValueError(f"Could not resolve normalization layer '{query}'")
34
+
35
+
36
+ def _activation_resolver(act, **kwargs):
37
+ """Resolves Keras and PyG-style activation names ("relu", "leaky_relu", "LeakyReLU", ...)."""
38
+ if act is None:
39
+ return None
40
+ if isinstance(act, str):
41
+ name = act.lower().replace("_", "")
42
+ aliases = {"leakyrelu": "leaky_relu", "hardswish": "hard_swish", "hardsigmoid": "hard_sigmoid"}
43
+ act = keras.activations.get(aliases.get(name, act.lower()))
44
+ if kwargs:
45
+ import functools
46
+
47
+ return functools.partial(act, **kwargs)
48
+ return act
49
+
50
+
51
+ class MLP(keras.Model):
52
+ r"""A Multi-Layer Perceptron (MLP) model.
53
+
54
+ There exists two ways to instantiate an `MLP`:
55
+
56
+ 1. By specifying explicit channel sizes, e.g., `MLP([16, 32, 64, 128])`
57
+ creates a three-layer MLP with **differently** sized hidden layers.
58
+
59
+ 2. By specifying fixed hidden channel sizes over a number of layers,
60
+ e.g., `MLP(in_channels=16, hidden_channels=32, out_channels=128,
61
+ num_layers=3)` creates a three-layer MLP with **equally** sized
62
+ hidden layers.
63
+
64
+ Args:
65
+ channel_list (List[int] or int, optional): List of input,
66
+ intermediate and output channels such that
67
+ `len(channel_list) - 1` denotes the number of layers of the
68
+ MLP. (default: `None`)
69
+ in_channels (int, optional): Size of each input sample. Will
70
+ override `channel_list`. (default: `None`)
71
+ hidden_channels (int, optional): Size of each hidden sample. Will
72
+ override `channel_list`. (default: `None`)
73
+ out_channels (int, optional): Size of each output sample. Will
74
+ override `channel_list`. (default: `None`)
75
+ num_layers (int, optional): The number of layers. Will override
76
+ `channel_list`. (default: `None`)
77
+ dropout (float or List[float], optional): Dropout probability of
78
+ each hidden embedding. (default: `0.`)
79
+ act (str or Callable, optional): The non-linear activation function
80
+ to use. (default: `"relu"`)
81
+ act_first (bool, optional): If set to `True`, activation is applied
82
+ before normalization. (default: `False`)
83
+ act_kwargs (dict, optional): Arguments passed to the activation function, e.g.
84
+ ``{"negative_slope": 0.2}`` for ``"leaky_relu"``. (default: `None`)
85
+ norm (str or Callable, optional): The normalization function to
86
+ use. (default: `"batch_norm"`)
87
+ norm_kwargs (dict, optional): Arguments passed to the respective
88
+ normalization function. (default: `None`)
89
+ plain_last (bool, optional): If set to `False`, will apply
90
+ non-linearity, normalization and dropout to the last layer as
91
+ well. (default: `True`)
92
+ bias (bool or List[bool], optional): If set to `False`, the module
93
+ will not learn additive biases. (default: `True`)
94
+
95
+ Example:
96
+ ```python
97
+ import numpy as np
98
+ from k3_node.models import MLP
99
+
100
+ x = np.random.rand(10, 16).astype("float32")
101
+ mlp = MLP([16, 32, 32, 4]) # channel sizes: input, hidden, hidden, output
102
+ print(tuple(mlp(x).shape)) # (10, 4)
103
+
104
+ mlp = MLP(in_channels=16, hidden_channels=32, out_channels=4, num_layers=3, dropout=0.1)
105
+ print(tuple(mlp(x).shape)) # (10, 4)
106
+ ```
107
+ """
108
+ def __init__(
109
+ self,
110
+ channel_list: Optional[Union[List[int], int]] = None,
111
+ *args,
112
+ in_channels: Optional[int] = None,
113
+ hidden_channels: Optional[int] = None,
114
+ out_channels: Optional[int] = None,
115
+ num_layers: Optional[int] = None,
116
+ dropout: Union[float, List[float]] = 0.0,
117
+ act="relu",
118
+ act_first: bool = False,
119
+ act_kwargs: Optional[dict] = None,
120
+ norm="batch_norm",
121
+ norm_kwargs: Optional[dict] = None,
122
+ plain_last: bool = True,
123
+ bias: Union[bool, List[bool]] = True,
124
+ **kwargs,
125
+ ):
126
+ super().__init__(**kwargs)
127
+
128
+ if len(args) > 0:
129
+ if isinstance(channel_list, int):
130
+ channel_list = [channel_list] + list(args)
131
+ elif isinstance(channel_list, (list, tuple)):
132
+ channel_list = list(channel_list) + list(args)
133
+
134
+ if isinstance(channel_list, int):
135
+ in_channels = channel_list
136
+ channel_list = None
137
+
138
+ if in_channels is not None:
139
+ if num_layers is None:
140
+ raise ValueError("Argument `num_layers` must be given")
141
+ if num_layers > 1 and hidden_channels is None:
142
+ raise ValueError(
143
+ f"Argument `hidden_channels` must be given for `num_layers={num_layers}`"
144
+ )
145
+ if out_channels is None:
146
+ raise ValueError("Argument `out_channels` must be given")
147
+
148
+ channel_list = [hidden_channels] * (num_layers - 1)
149
+ channel_list = [in_channels] + channel_list + [out_channels]
150
+
151
+ assert isinstance(channel_list, (tuple, list))
152
+ assert len(channel_list) >= 2
153
+ self.channel_list = list(channel_list)
154
+ self.in_channels = self.channel_list[0]
155
+ self.out_channels = self.channel_list[-1]
156
+
157
+ self.act = _activation_resolver(act, **(act_kwargs or {}))
158
+ self.act_first = act_first
159
+ self.plain_last = plain_last
160
+
161
+ if isinstance(dropout, float):
162
+ dropout = [dropout] * (len(channel_list) - 1)
163
+ if plain_last:
164
+ dropout[-1] = 0.0
165
+ if len(dropout) != len(channel_list) - 1:
166
+ raise ValueError(
167
+ f"Number of dropout values provided ({len(dropout)}) does not "
168
+ f"match the number of layers specified ({len(channel_list) - 1})"
169
+ )
170
+ self.dropout_rate = dropout
171
+
172
+ if isinstance(bias, bool):
173
+ bias = [bias] * (len(channel_list) - 1)
174
+ if len(bias) != len(channel_list) - 1:
175
+ raise ValueError(
176
+ f"Number of bias values provided ({len(bias)}) does not match "
177
+ f"the number of layers specified ({len(channel_list) - 1})"
178
+ )
179
+
180
+ self.lins = []
181
+ for in_c, out_c, _bias in zip(channel_list[:-1], channel_list[1:], bias):
182
+ lin = keras.layers.Dense(out_c, use_bias=_bias)
183
+ lin.build((None, in_c))
184
+ self.lins.append(lin)
185
+
186
+ self.norms = []
187
+ iterator = channel_list[1:-1] if plain_last else channel_list[1:]
188
+ for hc in iterator:
189
+ norm_layer = _normalization_resolver(norm, hc, **(norm_kwargs or {}))
190
+ if norm_layer is not None and hasattr(norm_layer, "build") and not norm_layer.built:
191
+ norm_layer.build((None, hc))
192
+ self.norms.append(norm_layer)
193
+
194
+ self.dropouts = [keras.layers.Dropout(p) if p > 0.0 else None for p in self.dropout_rate]
195
+
196
+ self.supports_norm_batch = False
197
+ if len(self.norms) > 0 and self.norms[0] is not None:
198
+ norm_params = inspect.signature(self.norms[0].call).parameters
199
+ self.supports_norm_batch = "batch" in norm_params
200
+ # Norms such as BatchNorm treat ``training=None`` as training mode, so pass it explicitly.
201
+ self.supports_norm_training = False
202
+ if len(self.norms) > 0 and self.norms[0] is not None:
203
+ self.supports_norm_training = "training" in inspect.signature(self.norms[0].call).parameters
204
+
205
+ @property
206
+ def num_layers(self) -> int:
207
+ r"""The number of layers."""
208
+ return len(self.channel_list) - 1
209
+
210
+ def build(self, input_shape=None):
211
+ for lin in self.lins:
212
+ if hasattr(lin, "built") and not lin.built:
213
+ lin.build(input_shape)
214
+ norm_iter = self.channel_list[1:-1] if self.plain_last else self.channel_list[1:]
215
+ for norm, hc in zip(self.norms, norm_iter):
216
+ if norm is not None and hasattr(norm, "built") and not norm.built:
217
+ norm.build((None, hc))
218
+ for drop in self.dropouts:
219
+ if drop is not None and hasattr(drop, "built") and not drop.built:
220
+ drop.build(input_shape)
221
+ self.built = True
222
+
223
+ def reset_parameters(self):
224
+ r"""Resets all learnable parameters of the module."""
225
+ for lin in self.lins:
226
+ if hasattr(lin, "kernel_initializer") and lin.kernel is not None:
227
+ lin.kernel.assign(lin.kernel_initializer(ops.shape(lin.kernel)))
228
+ if lin.bias is not None:
229
+ lin.bias.assign(lin.bias_initializer(ops.shape(lin.bias)))
230
+ for norm in self.norms:
231
+ if hasattr(norm, "reset_parameters"):
232
+ norm.reset_parameters()
233
+
234
+ def call(self, x, batch=None, batch_size=None, return_emb=None, training=None):
235
+ emb = None
236
+
237
+ # If `plain_last=True`, `len(norms) == len(lins) - 1`, thus skipping
238
+ # execution of the last layer inside the loop.
239
+ for i, (lin, norm) in enumerate(zip(self.lins, self.norms)):
240
+ x = lin(x)
241
+ if self.act is not None and self.act_first:
242
+ x = self.act(x)
243
+ if norm is not None:
244
+ norm_kwargs = {"training": training} if self.supports_norm_training else {}
245
+ if self.supports_norm_batch:
246
+ x = norm(x, batch, batch_size, **norm_kwargs)
247
+ else:
248
+ x = norm(x, **norm_kwargs)
249
+ if self.act is not None and not self.act_first:
250
+ x = self.act(x)
251
+ if self.dropouts[i] is not None:
252
+ x = self.dropouts[i](x, training=training)
253
+ if isinstance(return_emb, bool) and return_emb is True:
254
+ emb = x
255
+
256
+ if self.plain_last:
257
+ x = self.lins[-1](x)
258
+ if self.dropouts[-1] is not None:
259
+ x = self.dropouts[-1](x, training=training)
260
+
261
+ return (x, emb) if isinstance(return_emb, bool) else x
262
+
263
+ def __repr__(self) -> str:
264
+ return f"{self.__class__.__name__}({str(self.channel_list)[1:-1]})"