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.
- tennetsac-0.1.1/PKG-INFO +55 -0
- tennetsac-0.1.1/README.md +48 -0
- tennetsac-0.1.1/pyproject.toml +25 -0
- tennetsac-0.1.1/setup.cfg +4 -0
- tennetsac-0.1.1/tennetsac/__init__.py +1 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/base.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/1.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/10.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/2.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/3.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/4.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/5.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/6.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/7.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/8.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/fine-tuned/9.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/geo.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/ckpt_files/prf.ckpt +0 -0
- tennetsac-0.1.1/tennetsac/core.py +88 -0
- tennetsac-0.1.1/tennetsac/models/Emb2Geometry.py +105 -0
- tennetsac-0.1.1/tennetsac/models/Emb2Profile.py +85 -0
- tennetsac-0.1.1/tennetsac/models/Prf2Gamma.py +120 -0
- tennetsac-0.1.1/tennetsac/utils/embedding.py +50 -0
- tennetsac-0.1.1/tennetsac/utils/model_io.py +19 -0
- tennetsac-0.1.1/tennetsac/utils/plotting.py +82 -0
- tennetsac-0.1.1/tennetsac/utils/property.py +161 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/bert_vocab_curated.txt +2393 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/__init__.py +15 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/aggregate/__init__.py +128 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/__init__.py +20 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/attention_layer.py +113 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/causal_linear_attention.py +116 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/clustered_attention.py +195 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/conditional_full_attention.py +66 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/exact_topk_attention.py +88 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/full_attention.py +95 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/improved_clustered_attention.py +268 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/improved_clustered_causal_attention.py +257 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/linear_attention.py +92 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/local_attention.py +101 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention/reformer_attention.py +166 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/__init__.py +17 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/registry.py +61 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/spec.py +126 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/builders/__init__.py +59 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/builders/attention_builders.py +139 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/builders/base.py +67 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/builders/transformer_builders.py +550 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/causal_product/__init__.py +78 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/clustering/__init__.py +0 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/clustering/hamming/__init__.py +115 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/events/__init__.py +10 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/events/event.py +51 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/events/event_dispatcher.py +92 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/events/filters.py +141 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/__init__.py +12 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/base.py +73 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/fourier_features.py +287 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/hashing/__init__.py +31 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/local_product/__init__.py +97 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/masking.py +206 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/__init__.py +7 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/_utils.py +16 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/__init__.py +16 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/__init__.py +30 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/attention_layer.py +105 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/full_attention.py +75 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/linear_attention.py +79 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/__init__.py +30 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/attention_layer.py +96 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/full_attention.py +83 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/linear_attention.py +110 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/recurrent/transformers.py +279 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/sparse_product/__init__.py +399 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/transformers.py +294 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/utils.py +33 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/fast_transformers/weight_mapper.py +273 -0
- tennetsac-0.1.1/tennetsac/utils/smi_ted_light/load.py +680 -0
- tennetsac-0.1.1/tennetsac/utils/smiles.py +28 -0
- tennetsac-0.1.1/tennetsac.egg-info/PKG-INFO +55 -0
- tennetsac-0.1.1/tennetsac.egg-info/SOURCES.txt +156 -0
- tennetsac-0.1.1/tennetsac.egg-info/dependency_links.txt +1 -0
- tennetsac-0.1.1/tennetsac.egg-info/top_level.txt +1 -0
tennetsac-0.1.1/PKG-INFO
ADDED
|
@@ -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 @@
|
|
|
1
|
+
from .core import profile, binary_lng, multi_lng
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
@@ -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
|