tennetsac 0.1.1__tar.gz

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 (83) hide show
  1. tennetsac-0.1.1/PKG-INFO +55 -0
  2. tennetsac-0.1.1/README.md +48 -0
  3. tennetsac-0.1.1/pyproject.toml +25 -0
  4. tennetsac-0.1.1/setup.cfg +4 -0
  5. tennetsac-0.1.1/tennetsac/__init__.py +1 -0
  6. tennetsac-0.1.1/tennetsac/ckpt_files/base.ckpt +0 -0
  7. tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/1.ckpt +0 -0
  8. tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/10.ckpt +0 -0
  9. tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/2.ckpt +0 -0
  10. tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/3.ckpt +0 -0
  11. tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/4.ckpt +0 -0
  12. tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/5.ckpt +0 -0
  13. tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/6.ckpt +0 -0
  14. tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/7.ckpt +0 -0
  15. tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/8.ckpt +0 -0
  16. tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/9.ckpt +0 -0
  17. tennetsac-0.1.1/tennetsac/ckpt_files/geo.ckpt +0 -0
  18. tennetsac-0.1.1/tennetsac/ckpt_files/prf.ckpt +0 -0
  19. tennetsac-0.1.1/tennetsac/core.py +88 -0
  20. tennetsac-0.1.1/tennetsac/models/Emb2Geometry.py +105 -0
  21. tennetsac-0.1.1/tennetsac/models/Emb2Profile.py +85 -0
  22. tennetsac-0.1.1/tennetsac/models/Prf2Gamma.py +120 -0
  23. tennetsac-0.1.1/tennetsac/utils/embedding.py +50 -0
  24. tennetsac-0.1.1/tennetsac/utils/model_io.py +19 -0
  25. tennetsac-0.1.1/tennetsac/utils/plotting.py +82 -0
  26. tennetsac-0.1.1/tennetsac/utils/property.py +161 -0
  27. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/bert_vocab_curated.txt +2393 -0
  28. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/__init__.py +15 -0
  29. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/aggregate/__init__.py +128 -0
  30. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/__init__.py +20 -0
  31. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/attention_layer.py +113 -0
  32. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/causal_linear_attention.py +116 -0
  33. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/clustered_attention.py +195 -0
  34. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/conditional_full_attention.py +66 -0
  35. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/exact_topk_attention.py +88 -0
  36. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/full_attention.py +95 -0
  37. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/improved_clustered_attention.py +268 -0
  38. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/improved_clustered_causal_attention.py +257 -0
  39. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/linear_attention.py +92 -0
  40. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/local_attention.py +101 -0
  41. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/reformer_attention.py +166 -0
  42. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/__init__.py +17 -0
  43. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/registry.py +61 -0
  44. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/spec.py +126 -0
  45. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/builders/__init__.py +59 -0
  46. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/builders/attention_builders.py +139 -0
  47. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/builders/base.py +67 -0
  48. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/builders/transformer_builders.py +550 -0
  49. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/causal_product/__init__.py +78 -0
  50. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/clustering/__init__.py +0 -0
  51. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/clustering/hamming/__init__.py +115 -0
  52. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/events/__init__.py +10 -0
  53. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/events/event.py +51 -0
  54. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/events/event_dispatcher.py +92 -0
  55. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/events/filters.py +141 -0
  56. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/__init__.py +12 -0
  57. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/base.py +73 -0
  58. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/fourier_features.py +287 -0
  59. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/hashing/__init__.py +31 -0
  60. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/local_product/__init__.py +97 -0
  61. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/masking.py +206 -0
  62. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/__init__.py +7 -0
  63. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/_utils.py +16 -0
  64. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/__init__.py +16 -0
  65. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/__init__.py +30 -0
  66. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/attention_layer.py +105 -0
  67. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/full_attention.py +75 -0
  68. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/linear_attention.py +79 -0
  69. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/__init__.py +30 -0
  70. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/attention_layer.py +96 -0
  71. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/full_attention.py +83 -0
  72. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/linear_attention.py +110 -0
  73. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/transformers.py +279 -0
  74. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/sparse_product/__init__.py +399 -0
  75. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/transformers.py +294 -0
  76. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/utils.py +33 -0
  77. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/weight_mapper.py +273 -0
  78. tennetsac-0.1.1/tennetsac/utils/smi_ted_light/load.py +680 -0
  79. tennetsac-0.1.1/tennetsac/utils/smiles.py +28 -0
  80. tennetsac-0.1.1/tennetsac.egg-info/PKG-INFO +55 -0
  81. tennetsac-0.1.1/tennetsac.egg-info/SOURCES.txt +156 -0
  82. tennetsac-0.1.1/tennetsac.egg-info/dependency_links.txt +1 -0
  83. tennetsac-0.1.1/tennetsac.egg-info/top_level.txt +1 -0
@@ -0,0 +1,55 @@
1
+ Metadata-Version: 2.4
2
+ Name: tennetsac
3
+ Version: 0.1.1
4
+ Summary: Thermodynamics-embedded Neural Network for Segment Activity Coefficients
5
+ Author: Yue Yang
6
+ Description-Content-Type: text/markdown
7
+
8
+ # TeNNet-SAC
9
+
10
+ TeNNet-SAC (Thermodynamics-Embedded Neural Network for Segment Activity Coefficients) is a machine learning framework designed to predict molecular activity coefficients in multicomponent systems using only molecular SMILES strings, composition, and temperature as input.
11
+
12
+ This project provides:
13
+
14
+ 1. **σ-profile prediction model**, including surface area and molecular volume estimation
15
+ 2. **Activity coefficient prediction** with two versions:
16
+ - **Base model**: trained on synthetic data generated from the COSMO-SAC model
17
+ - **Fine-tuned model**: further optimized using high-quality experimental data
18
+
19
+ ## Features
20
+
21
+ - Predicts activity coefficients beyond binary systems
22
+ - Requires only SMILES strings, mole fractions, and temperature
23
+ - Hard-constraint architecture ensures thermodynamic consistency
24
+ - Modular design: supports use of σ-profiles from QC calculations
25
+ - Robust and generalizable via two-stage training: synthetic COSMO-SAC pretraining followed by experimental fine-tuning, preserving physical consistency
26
+
27
+ ## Citation
28
+
29
+ Yue Yang, Shiang-Tai Lin. *Physics-Embedded Machine Learning Model for Phase Equilibrium Prediction in Multicomponent Systems*. *Journal of Chemical Information and Modeling*, 2025. [DOI: 10.1021/acs.jcim.5c01804](https://doi.org/10.1021/acs.jcim.5c01804)
30
+
31
+ ## References
32
+
33
+ This project builds upon the following foundational models. If you use this project in your research, we encourage you to cite them as well:
34
+
35
+ - **ChemBERTa-2**
36
+ Ahmad, W.; Simon, E.; Chithrananda, S.; Grand, G.; Ramsundar, B. Chemberta-2: Towards chemical foundation models. arXiv preprint arXiv:2209.01712 2022.
37
+ [https://arxiv.org/abs/2209.01712](https://arxiv.org/abs/2209.01712)
38
+
39
+ - **SMI-TED**
40
+ Soares, E.; Shirasuna, V.; Brazil, E. V.; Cerqueira, R.; Zubarev, D.; Schmidt, K. A large encoder-decoder family of foundation models for chemical language. arXiv preprint arXiv:2407.20267 2024.
41
+ [https://arxiv.org/abs/2407.20267](https://arxiv.org/abs/2407.20267)
42
+
43
+ ### External Code Acknowledgment
44
+
45
+ The folder `smi_ted_light/` is adapted from the [SMI-TED](https://github.com/IBM/materials/tree/main/models/smi_ted) repository by Soares et al., with only minimal modifications. The core implementation remains unchanged. Full credit goes to the original authors.
46
+
47
+ ## License
48
+
49
+ MIT License. See [LICENSE](LICENSE) for details.
50
+
51
+ ---
52
+
53
+ Maintained by **Yue Yang** ([@yueyue2299](https://github.com/yueyue2299)).
54
+
55
+ COMET, Department of Chemical Engineering, National Taiwan University
@@ -0,0 +1,48 @@
1
+ # TeNNet-SAC
2
+
3
+ TeNNet-SAC (Thermodynamics-Embedded Neural Network for Segment Activity Coefficients) is a machine learning framework designed to predict molecular activity coefficients in multicomponent systems using only molecular SMILES strings, composition, and temperature as input.
4
+
5
+ This project provides:
6
+
7
+ 1. **σ-profile prediction model**, including surface area and molecular volume estimation
8
+ 2. **Activity coefficient prediction** with two versions:
9
+ - **Base model**: trained on synthetic data generated from the COSMO-SAC model
10
+ - **Fine-tuned model**: further optimized using high-quality experimental data
11
+
12
+ ## Features
13
+
14
+ - Predicts activity coefficients beyond binary systems
15
+ - Requires only SMILES strings, mole fractions, and temperature
16
+ - Hard-constraint architecture ensures thermodynamic consistency
17
+ - Modular design: supports use of σ-profiles from QC calculations
18
+ - Robust and generalizable via two-stage training: synthetic COSMO-SAC pretraining followed by experimental fine-tuning, preserving physical consistency
19
+
20
+ ## Citation
21
+
22
+ Yue Yang, Shiang-Tai Lin. *Physics-Embedded Machine Learning Model for Phase Equilibrium Prediction in Multicomponent Systems*. *Journal of Chemical Information and Modeling*, 2025. [DOI: 10.1021/acs.jcim.5c01804](https://doi.org/10.1021/acs.jcim.5c01804)
23
+
24
+ ## References
25
+
26
+ This project builds upon the following foundational models. If you use this project in your research, we encourage you to cite them as well:
27
+
28
+ - **ChemBERTa-2**
29
+ Ahmad, W.; Simon, E.; Chithrananda, S.; Grand, G.; Ramsundar, B. Chemberta-2: Towards chemical foundation models. arXiv preprint arXiv:2209.01712 2022.
30
+ [https://arxiv.org/abs/2209.01712](https://arxiv.org/abs/2209.01712)
31
+
32
+ - **SMI-TED**
33
+ Soares, E.; Shirasuna, V.; Brazil, E. V.; Cerqueira, R.; Zubarev, D.; Schmidt, K. A large encoder-decoder family of foundation models for chemical language. arXiv preprint arXiv:2407.20267 2024.
34
+ [https://arxiv.org/abs/2407.20267](https://arxiv.org/abs/2407.20267)
35
+
36
+ ### External Code Acknowledgment
37
+
38
+ The folder `smi_ted_light/` is adapted from the [SMI-TED](https://github.com/IBM/materials/tree/main/models/smi_ted) repository by Soares et al., with only minimal modifications. The core implementation remains unchanged. Full credit goes to the original authors.
39
+
40
+ ## License
41
+
42
+ MIT License. See [LICENSE](LICENSE) for details.
43
+
44
+ ---
45
+
46
+ Maintained by **Yue Yang** ([@yueyue2299](https://github.com/yueyue2299)).
47
+
48
+ COMET, Department of Chemical Engineering, National Taiwan University
@@ -0,0 +1,25 @@
1
+ [build-system]
2
+ requires = ["setuptools>=61.0"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [tool.setuptools]
6
+ package-dir = {"" = "."}
7
+
8
+ [tool.setuptools.packages.find]
9
+ include = ["tennetsac*"]
10
+
11
+ [tool.setuptools.package-data]
12
+ tennetsac = [
13
+ "utils/smi_ted_light/*.txt",
14
+ "ckpt_files/*.ckpt",
15
+ "ckpt_files/fine-tuned/*.ckpt"
16
+ ]
17
+
18
+ [project]
19
+ name = "tennetsac"
20
+ version = "0.1.1"
21
+ description = "Thermodynamics-embedded Neural Network for Segment Activity Coefficients"
22
+ authors = [{ name="Yue Yang" }]
23
+ readme = "README.md"
24
+ license = { file = "MIT" }
25
+
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+
@@ -0,0 +1 @@
1
+ from .core import profile, binary_lng, multi_lng
@@ -0,0 +1,88 @@
1
+ # === Standard Library ===
2
+ import os
3
+ import warnings
4
+
5
+ # === Computing & Visualization ===
6
+ import torch
7
+ import pandas as pd
8
+ import numpy as np
9
+
10
+ # === Warning Suppression ===
11
+ from tqdm import TqdmWarning
12
+ warnings.simplefilter("ignore", category=TqdmWarning)
13
+ warnings.filterwarnings("ignore", message="TypedStorage is deprecated")
14
+ warnings.filterwarnings("ignore", category=FutureWarning)
15
+ warnings.filterwarnings("ignore", message="Some weights of RobertaModel were not initialized")
16
+ from transformers.utils import logging
17
+ logging.set_verbosity_error()
18
+
19
+ # === Model Architectures ===
20
+ from .models.Emb2Profile import SigmaProfileGenerator
21
+ from .models.Emb2Geometry import GeometryGenerator
22
+ from .models.Prf2Gamma import Prf_to_Seg_Model
23
+
24
+ # === Model Loading ===
25
+ from .utils.model_io import load_model, load_all_Gamma_models
26
+
27
+ # === Embedding Extraction ===
28
+
29
+ from .utils.embedding import ChemBERTaEmbedder, SMITEDEmbedder
30
+
31
+ # === Computation ===
32
+ from .utils.property import get_sigma_profile, calc_ln_gamma, ensemble_segac, calc_ln_gamma_binary
33
+
34
+ # === Embedding models ===
35
+ cb_emb = ChemBERTaEmbedder()
36
+ st_emb = SMITEDEmbedder()
37
+
38
+ # === Load checkpoints ===
39
+ here = os.path.dirname(__file__)
40
+ ckpt_path = os.path.join(here, "ckpt_files")
41
+
42
+ prf_model = load_model(SigmaProfileGenerator(), os.path.join(ckpt_path, "prf.ckpt"))
43
+ geometry_model = load_model(GeometryGenerator(), os.path.join(ckpt_path, "geo.ckpt"))
44
+ Gamma_base_model = load_model(Prf_to_Seg_Model(), os.path.join(ckpt_path, "base.ckpt"))
45
+
46
+ Gamma_finetuned_models = load_all_Gamma_models(Prf_to_Seg_Model, os.path.join(ckpt_path, "fine-tuned"))
47
+
48
+ # === Define functions ===
49
+ def sigma_profile_wrapper(smiles):
50
+ return get_sigma_profile(smiles, prf_model, geometry_model, cb_emb, st_emb)
51
+
52
+ def single_model_predictor(sigma, temperature):
53
+ return Gamma_base_model(sigma, torch.tensor([temperature]))[1]
54
+
55
+ def ensemble_predictor(sigma, temperature):
56
+ return ensemble_segac(Gamma_finetuned_models, sigma, temperature)
57
+
58
+ def select_gamma_predictor(model_type: str):
59
+ if model_type == "base":
60
+ return single_model_predictor
61
+ elif model_type == "tuned":
62
+ return ensemble_predictor
63
+ else:
64
+ raise ValueError("Invalid model_type. Choose 'base' or 'tuned'.")
65
+
66
+ def profile(smiles):
67
+ return sigma_profile_wrapper(smiles)
68
+
69
+ def binary_lng(smiles:list, temperature:float, molefraction:list):
70
+ gamma_predictor = select_gamma_predictor("tuned")
71
+
72
+ ln_gamma_1, ln_gamma_2 = calc_ln_gamma_binary(smiles[0], smiles[1], molefraction, temperature,
73
+ gamma_predictor=gamma_predictor,
74
+ get_sigma_profile_fn=sigma_profile_wrapper)
75
+ return ln_gamma_1.tolist(), ln_gamma_2.tolist()
76
+
77
+ def multi_lng(smiles:list, temperature:float, composition:list):
78
+ gamma_predictor = select_gamma_predictor("tuned")
79
+
80
+ lng_array = calc_ln_gamma(
81
+ smiles,
82
+ composition,
83
+ temperature,
84
+ gamma_predictor=gamma_predictor,
85
+ get_sigma_profile_fn=sigma_profile_wrapper
86
+ )
87
+
88
+ return lng_array
@@ -0,0 +1,105 @@
1
+ import torch
2
+ import torch.nn as nn
3
+
4
+ class ResidualBlock(nn.Module): # ResNet v2
5
+ def __init__(self, in_features, hidden_features=None, activation=nn.GELU):
6
+ """
7
+ in_features: input feature
8
+ hidden_features: hidden feature, if None, use in_features
9
+ """
10
+ super(ResidualBlock, self).__init__()
11
+ hidden_features = hidden_features or in_features
12
+ self.act = activation()
13
+ self.bn1 = nn.BatchNorm1d(hidden_features)
14
+ self.fc1 = nn.Linear(in_features, hidden_features)
15
+ self.bn2 = nn.BatchNorm1d(in_features)
16
+ self.fc2 = nn.Linear(hidden_features, in_features)
17
+
18
+ def forward(self, x):
19
+ identity = x
20
+
21
+ out = self.bn1(x)
22
+ out = self.act(out)
23
+ out = self.fc1(out)
24
+
25
+ out = self.bn2(out)
26
+ out = self.act(out)
27
+ out = self.fc2(out)
28
+
29
+ out += identity # add residual
30
+ return out
31
+
32
+ class GeometryGenerator(nn.Module):
33
+ def __init__(self, STemb_dim=768, CBemb_dim=384, CBout_dim=32, av_dim=2,
34
+ hidden_dim=128, mw_dim=128, dropout_rate=0.1):
35
+ """
36
+ sig_dim: int, dimension of sigma profile input
37
+ """
38
+ super(GeometryGenerator, self).__init__()
39
+
40
+ self.isomerism_encoder = nn.Sequential(
41
+ nn.Linear(CBemb_dim, 256),
42
+ nn.GELU(),
43
+ nn.Linear(256, CBout_dim)
44
+ )
45
+
46
+ self.backbone = nn.Sequential(
47
+ nn.Linear(STemb_dim + CBout_dim, 512),
48
+ nn.GELU(),
49
+ nn.Linear(512, hidden_dim),
50
+ nn.GELU(),
51
+ )
52
+
53
+ self.mw_projector = nn.Sequential(
54
+ nn.Linear(1, 64),
55
+ nn.GELU(),
56
+ nn.Linear(64, mw_dim),
57
+ nn.GELU(),
58
+ )
59
+
60
+ self.comb_layers = nn.Sequential(
61
+ nn.Linear(hidden_dim + mw_dim, hidden_dim),
62
+ nn.GELU(),
63
+ nn.Linear(hidden_dim, hidden_dim),
64
+ nn.GELU(),
65
+ )
66
+
67
+ self.res_block = ResidualBlock(in_features=hidden_dim, hidden_features=hidden_dim, activation=nn.GELU)
68
+
69
+ self.area_head = nn.Sequential(
70
+ nn.Linear(hidden_dim, 64),
71
+ nn.GELU(),
72
+ # nn.Dropout(p=0.1),
73
+ nn.Linear(64, 16),
74
+ nn.GELU(),
75
+ nn.Linear(16, 1),
76
+ )
77
+
78
+ self.volume_head = nn.Sequential(
79
+ nn.Linear(hidden_dim, 64),
80
+ nn.GELU(),
81
+ # nn.Dropout(p=0.1),
82
+ nn.Linear(64, 16),
83
+ nn.GELU(),
84
+ nn.Linear(16, 1),
85
+ )
86
+
87
+ def forward(self, STemb, CBemb, mw):
88
+ mw = mw.unsqueeze(1)
89
+ conf_feat = self.isomerism_encoder(CBemb)
90
+
91
+ emb = torch.cat([STemb, conf_feat], dim=1)
92
+
93
+ hidden_emb = self.backbone(emb)
94
+ hidden_mw = self.mw_projector(mw)
95
+
96
+ hidden_comb = torch.cat([hidden_emb, hidden_mw], dim=1)
97
+ hidden = self.comb_layers(hidden_comb)
98
+ hidden = self.res_block(hidden)
99
+
100
+ area = self.area_head(hidden)
101
+ volume = self.volume_head(hidden)
102
+
103
+ output = torch.cat([area, volume], dim=1)
104
+
105
+ return output
@@ -0,0 +1,85 @@
1
+ import torch
2
+ import torch.nn as nn
3
+
4
+ class ResidualBlock(nn.Module): # ResNet v2
5
+ def __init__(self, in_features, hidden_features=None, activation=nn.GELU):
6
+ """
7
+ in_features: input feature
8
+ hidden_features: hidden feature, if None, use in_features
9
+ """
10
+ super(ResidualBlock, self).__init__()
11
+ hidden_features = hidden_features or in_features
12
+ self.act = activation()
13
+ self.bn1 = nn.BatchNorm1d(hidden_features)
14
+ self.fc1 = nn.Linear(in_features, hidden_features)
15
+ self.bn2 = nn.BatchNorm1d(in_features)
16
+ self.fc2 = nn.Linear(hidden_features, in_features)
17
+
18
+ def forward(self, x):
19
+ identity = x
20
+
21
+ out = self.bn1(x)
22
+ out = self.act(out)
23
+ out = self.fc1(out)
24
+
25
+ out = self.bn2(out)
26
+ out = self.act(out)
27
+ out = self.fc2(out)
28
+
29
+ out += identity # add residual
30
+ return out
31
+
32
+ class SigmaProfileGenerator(nn.Module):
33
+ def __init__(self, STemb_dim=768, CBemb_dim=384, CBout_dim=32,
34
+ prf_dim=51, hidden_dim=256, dropout_rate=0.3):
35
+ """
36
+ sig_dim: int, dimension of sigma profile input
37
+ """
38
+ super(SigmaProfileGenerator, self).__init__()
39
+
40
+ self.isomerism_encoder = nn.Sequential(
41
+ nn.Linear(CBemb_dim, 256),
42
+ nn.GELU(),
43
+ nn.Linear(256, CBout_dim)
44
+ )
45
+
46
+ self.backbone = nn.Sequential(
47
+ nn.Linear(STemb_dim + CBout_dim, 512),
48
+ nn.GELU(),
49
+ nn.Linear(512, 256),
50
+ nn.GELU(),
51
+ nn.Linear(256, hidden_dim),
52
+ nn.GELU(),
53
+ ResidualBlock(in_features=hidden_dim, hidden_features=hidden_dim, activation=nn.GELU),
54
+ nn.GELU(),
55
+ ResidualBlock(in_features=hidden_dim, hidden_features=hidden_dim, activation=nn.GELU),
56
+ nn.GELU(),
57
+ nn.Linear(hidden_dim, hidden_dim),
58
+ nn.GELU()
59
+ )
60
+
61
+ self.head = nn.Sequential(
62
+ nn.Linear(hidden_dim, 128),
63
+ nn.GELU(),
64
+ nn.Dropout(p=dropout_rate),
65
+ nn.Linear(128, 128),
66
+ nn.GELU(),
67
+ nn.Dropout(p=dropout_rate),
68
+ nn.Linear(128, 128),
69
+ nn.GELU(),
70
+ nn.Dropout(p=dropout_rate),
71
+ nn.Linear(128, 64),
72
+ nn.GELU(),
73
+ nn.Linear(64, prf_dim),
74
+ nn.ReLU()
75
+ )
76
+
77
+ def forward(self, STemb, CBemb):
78
+ conf_feat = self.isomerism_encoder(CBemb)
79
+ x = torch.cat([STemb, conf_feat], dim=1)
80
+ x = self.backbone(x)
81
+ output = self.head(x)
82
+ area = output.sum(dim=1, keepdim=True)
83
+ prf = output / area
84
+
85
+ return prf
@@ -0,0 +1,120 @@
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.autograd as autograd
4
+
5
+ class ResidualBlock(nn.Module):
6
+ def __init__(self, in_features, hidden_features=None, activation=nn.GELU):
7
+ """
8
+ in_features: input feature
9
+ hidden_features: hidden feature, if None, use in_features
10
+ """
11
+ super(ResidualBlock, self).__init__()
12
+ hidden_features = hidden_features or in_features
13
+ self.fc1 = nn.Linear(in_features, hidden_features)
14
+ self.act = activation()
15
+ self.fc2 = nn.Linear(hidden_features, in_features)
16
+
17
+ def forward(self, x):
18
+ identity = x
19
+ out = self.fc1(x)
20
+ out = self.act(out)
21
+ out = self.fc2(out)
22
+ out += identity # add residual
23
+ out = self.act(out)
24
+ return out
25
+
26
+ class Prf_to_Seg_Model(nn.Module):
27
+ def __init__(self, sig_dim=51, temp_hidden_dim=16, sig_hidden_dim=512, hidden_2_dim=256):
28
+ """
29
+ sig_dim: int, dimension of sigma profile input
30
+ """
31
+ super(Prf_to_Seg_Model, self).__init__()
32
+ # sigma model
33
+ self.model_sigma = nn.Sequential(
34
+ nn.Linear(sig_dim, 512),
35
+ nn.GELU(),
36
+ nn.Linear(512, 512),
37
+ nn.GELU(),
38
+ nn.Linear(512, sig_hidden_dim),
39
+ nn.GELU(),
40
+ )
41
+
42
+ self.bn_sig = nn.BatchNorm1d(num_features=sig_hidden_dim)
43
+
44
+ self.temp_embedding = nn.Sequential(
45
+ nn.Linear(1, 32),
46
+ nn.ReLU(),
47
+ nn.Linear(32, temp_hidden_dim),
48
+ nn.ReLU()
49
+ )
50
+
51
+ self.bn_t = nn.BatchNorm1d(num_features=temp_hidden_dim)
52
+
53
+ self.model_combined = nn.Sequential(
54
+ nn.Linear(sig_hidden_dim + temp_hidden_dim, 256),
55
+ nn.GELU(),
56
+ nn.Linear(256, hidden_2_dim),
57
+ nn.GELU(),
58
+ )
59
+
60
+ self.bn2 = nn.BatchNorm1d(num_features=hidden_2_dim)
61
+
62
+ # Resudual Block
63
+ self.res_block = ResidualBlock(in_features=hidden_2_dim, hidden_features=hidden_2_dim, activation=nn.GELU)
64
+
65
+ self.model_final = nn.Sequential(
66
+ nn.Linear(hidden_2_dim, 128),
67
+ nn.GELU(),
68
+ nn.Linear(128, 64),
69
+ nn.GELU(),
70
+ nn.Linear(64, 1),
71
+ )
72
+
73
+ def forward(self, sigs, t):
74
+ """
75
+ sigs: (batch_size, sig_dim)
76
+ t : (batch_size,) or (batch_size, 1)
77
+ """
78
+ # Ensure sigs require gradients
79
+ sigs = sigs.requires_grad_(True)
80
+ # sigs.requires_grad = True
81
+ t = t.requires_grad_(False)
82
+
83
+ # sigma profile
84
+ sum_sigs = sigs.sum(dim=1, keepdim=True) # (batch_size, 1)
85
+ # Send normalized profile into sigma model.
86
+ sigma_emb = self.model_sigma(sigs/sum_sigs) # (batch_size, hidden_1_dim)
87
+ sigma_emb = self.bn_sig(sigma_emb)
88
+
89
+ # temperature
90
+ t_inv = 1.0 / t.view(-1, 1) # (batch_size, 1)
91
+ t_emb = self.temp_embedding(t_inv) # (batch_size, temp_hidden_dim)
92
+ t_emb = self.bn_t(t_emb)
93
+
94
+ # Concatenate sigma_emb and t_emb
95
+ combined_in = torch.cat([sigma_emb, t_emb], dim=-1)
96
+ combined_out = self.model_combined(combined_in)
97
+
98
+ # Batch Normal 2
99
+ combined_out = self.bn2(combined_out)
100
+
101
+ # Residual connection
102
+ combined_out = self.res_block(combined_out)
103
+
104
+ # combined_out_t = torch.cat([combined_out, t_inv], dim=-1)
105
+
106
+ # Final FNN
107
+ gchg_part = self.model_final(combined_out)
108
+
109
+ # gchg = sum(sigs) * gchg_part
110
+ gchg = sum_sigs * gchg_part # (batch_size, 1)
111
+
112
+ # Summing gchg and compute the gradient
113
+ grad = autograd.grad(
114
+ gchg.sum(), # -> scalar
115
+ sigs,
116
+ create_graph=self.training
117
+ )[0] # shape: (batch_size, sig_dim)
118
+ segs = grad
119
+ return gchg, segs
120
+
@@ -0,0 +1,50 @@
1
+ import os
2
+ from transformers import RobertaTokenizer, RobertaModel
3
+ from .smiles import canonicalize_smiles, compute_mol_weight
4
+ import torch
5
+
6
+ # === ChemBERTa2 Embedding ===
7
+ class ChemBERTaEmbedder:
8
+ def __init__(self, model_name="DeepChem/ChemBERTa-77M-MLM", max_length=128, device="cpu"):
9
+ self.tokenizer = RobertaTokenizer.from_pretrained(model_name)
10
+ self.model = RobertaModel.from_pretrained(model_name).to(device)
11
+ self.max_length = max_length
12
+ self.device = device
13
+
14
+ def __call__(self, smiles):
15
+ tokens = self.tokenizer(
16
+ smiles, return_tensors="pt", padding="max_length",
17
+ truncation=True, max_length=self.max_length
18
+ ).to(self.device)
19
+ with torch.no_grad():
20
+ output = self.model(tokens["input_ids"], tokens["attention_mask"])
21
+ hidden = output.last_hidden_state
22
+ mask = tokens["attention_mask"].unsqueeze(-1)
23
+ avg_emb = (hidden * mask).sum(dim=1) / mask.sum(dim=1)
24
+ return avg_emb
25
+
26
+ # === SMI-TED Embedding ===
27
+ from .smi_ted_light.load import load_smi_ted
28
+
29
+ class SMITEDEmbedder:
30
+ def __init__(self, ckpt_name="smi-ted-Light_40.pt", device="cpu"):
31
+ here = os.path.dirname(__file__)
32
+ model_dir = os.path.join(here, "smi_ted_light")
33
+
34
+ self.model = load_smi_ted(folder=model_dir, ckpt_filename=ckpt_name).to(device)
35
+ self.device = device
36
+
37
+ def __call__(self, smiles):
38
+ with torch.no_grad():
39
+ emb = self.model.encode(smiles, return_torch=True).to("cpu") # always return CPU tensor
40
+ return emb
41
+
42
+ # === function ===
43
+ def get_input_embeddings(smiles, cb_embedder, st_embedder):
44
+ isomeric_cansmi = canonicalize_smiles(smiles, isomeric=True)
45
+ nonisomeric_cansmi = canonicalize_smiles(smiles, isomeric=False)
46
+
47
+ cb_emb = cb_embedder(isomeric_cansmi)
48
+ st_emb = st_embedder(nonisomeric_cansmi)
49
+
50
+ return st_emb, cb_emb
@@ -0,0 +1,19 @@
1
+ import torch
2
+ from pathlib import Path
3
+
4
+ def load_model(model, ckpt_file, device='cpu'):
5
+ model.to(torch.device(device))
6
+ model.load_state_dict(torch.load(ckpt_file, map_location=torch.device(device)))
7
+
8
+ return model.eval()
9
+
10
+ def load_all_Gamma_models(model_class, ckpt_dir, num_models=10):
11
+ models = []
12
+ for i in range(1, num_models + 1):
13
+ ckpt_file = Path(ckpt_dir) / f"{i}.ckpt"
14
+ model = model_class()
15
+ model.to(torch.device('cpu'))
16
+ model.load_state_dict(torch.load(ckpt_file, map_location=torch.device('cpu')))
17
+ model.eval()
18
+ models.append(model)
19
+ return models