tennetsac 0.1.1__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.
- tennetsac/__init__.py +1 -0
- tennetsac/ckpt_files/base.ckpt +0 -0
- tennetsac/ckpt_files/fine-tuned/1.ckpt +0 -0
- tennetsac/ckpt_files/fine-tuned/10.ckpt +0 -0
- tennetsac/ckpt_files/fine-tuned/2.ckpt +0 -0
- tennetsac/ckpt_files/fine-tuned/3.ckpt +0 -0
- tennetsac/ckpt_files/fine-tuned/4.ckpt +0 -0
- tennetsac/ckpt_files/fine-tuned/5.ckpt +0 -0
- tennetsac/ckpt_files/fine-tuned/6.ckpt +0 -0
- tennetsac/ckpt_files/fine-tuned/7.ckpt +0 -0
- tennetsac/ckpt_files/fine-tuned/8.ckpt +0 -0
- tennetsac/ckpt_files/fine-tuned/9.ckpt +0 -0
- tennetsac/ckpt_files/geo.ckpt +0 -0
- tennetsac/ckpt_files/prf.ckpt +0 -0
- tennetsac/core.py +88 -0
- tennetsac/models/Emb2Geometry.py +105 -0
- tennetsac/models/Emb2Profile.py +85 -0
- tennetsac/models/Prf2Gamma.py +120 -0
- tennetsac/utils/embedding.py +50 -0
- tennetsac/utils/model_io.py +19 -0
- tennetsac/utils/plotting.py +82 -0
- tennetsac/utils/property.py +161 -0
- tennetsac/utils/smi_ted_light/bert_vocab_curated.txt +2393 -0
- tennetsac/utils/smi_ted_light/fast_transformers/__init__.py +15 -0
- tennetsac/utils/smi_ted_light/fast_transformers/aggregate/__init__.py +128 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/__init__.py +20 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/attention_layer.py +113 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/causal_linear_attention.py +116 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/clustered_attention.py +195 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/conditional_full_attention.py +66 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/exact_topk_attention.py +88 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/full_attention.py +95 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/improved_clustered_attention.py +268 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/improved_clustered_causal_attention.py +257 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/linear_attention.py +92 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/local_attention.py +101 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention/reformer_attention.py +166 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/__init__.py +17 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/registry.py +61 -0
- tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/spec.py +126 -0
- tennetsac/utils/smi_ted_light/fast_transformers/builders/__init__.py +59 -0
- tennetsac/utils/smi_ted_light/fast_transformers/builders/attention_builders.py +139 -0
- tennetsac/utils/smi_ted_light/fast_transformers/builders/base.py +67 -0
- tennetsac/utils/smi_ted_light/fast_transformers/builders/transformer_builders.py +550 -0
- tennetsac/utils/smi_ted_light/fast_transformers/causal_product/__init__.py +78 -0
- tennetsac/utils/smi_ted_light/fast_transformers/clustering/__init__.py +0 -0
- tennetsac/utils/smi_ted_light/fast_transformers/clustering/hamming/__init__.py +115 -0
- tennetsac/utils/smi_ted_light/fast_transformers/events/__init__.py +10 -0
- tennetsac/utils/smi_ted_light/fast_transformers/events/event.py +51 -0
- tennetsac/utils/smi_ted_light/fast_transformers/events/event_dispatcher.py +92 -0
- tennetsac/utils/smi_ted_light/fast_transformers/events/filters.py +141 -0
- tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/__init__.py +12 -0
- tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/base.py +73 -0
- tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/fourier_features.py +287 -0
- tennetsac/utils/smi_ted_light/fast_transformers/hashing/__init__.py +31 -0
- tennetsac/utils/smi_ted_light/fast_transformers/local_product/__init__.py +97 -0
- tennetsac/utils/smi_ted_light/fast_transformers/masking.py +206 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/__init__.py +7 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/_utils.py +16 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/__init__.py +16 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/__init__.py +30 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/attention_layer.py +105 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/full_attention.py +75 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/linear_attention.py +79 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/__init__.py +30 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/attention_layer.py +96 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/full_attention.py +83 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/linear_attention.py +110 -0
- tennetsac/utils/smi_ted_light/fast_transformers/recurrent/transformers.py +279 -0
- tennetsac/utils/smi_ted_light/fast_transformers/sparse_product/__init__.py +399 -0
- tennetsac/utils/smi_ted_light/fast_transformers/transformers.py +294 -0
- tennetsac/utils/smi_ted_light/fast_transformers/utils.py +33 -0
- tennetsac/utils/smi_ted_light/fast_transformers/weight_mapper.py +273 -0
- tennetsac/utils/smi_ted_light/load.py +680 -0
- tennetsac/utils/smiles.py +28 -0
- tennetsac-0.1.1.dist-info/METADATA +55 -0
- tennetsac-0.1.1.dist-info/RECORD +79 -0
- tennetsac-0.1.1.dist-info/WHEEL +5 -0
- tennetsac-0.1.1.dist-info/top_level.txt +1 -0
tennetsac/__init__.py
ADDED
|
@@ -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
|
tennetsac/core.py
ADDED
|
@@ -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
|
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
import matplotlib.pyplot as plt
|
|
2
|
+
import matplotlib as mpl
|
|
3
|
+
import numpy as np
|
|
4
|
+
import pandas as pd
|
|
5
|
+
mpl.rcParams.update(mpl.rcParamsDefault)
|
|
6
|
+
|
|
7
|
+
def plot_sigma_prf(smiles, prf_tensor, savepath=None):
|
|
8
|
+
prf = prf_tensor.squeeze().detach().cpu().numpy()
|
|
9
|
+
sigma = np.linspace(-0.025, 0.025, 51)
|
|
10
|
+
plt.figure(figsize=(6, 4))
|
|
11
|
+
plt.plot(sigma, prf, linestyle='-', linewidth=2)
|
|
12
|
+
|
|
13
|
+
plt.xlim(-0.025, 0.025)
|
|
14
|
+
plt.xlabel("σ (e/Ų)", fontsize=12)
|
|
15
|
+
plt.ylabel("P(σ)", fontsize=12)
|
|
16
|
+
plt.title(f"σ-Profile: {smiles}", fontsize=13)
|
|
17
|
+
|
|
18
|
+
plt.xticks(fontsize=10)
|
|
19
|
+
plt.yticks(fontsize=10)
|
|
20
|
+
|
|
21
|
+
plt.tight_layout()
|
|
22
|
+
|
|
23
|
+
if savepath:
|
|
24
|
+
plt.savefig(savepath, dpi=300, bbox_inches="tight")
|
|
25
|
+
plt.show()
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def plot_sigma_prf_pair(smiles_1, prf_tensor_1, smiles_2, prf_tensor_2, savepath=None):
|
|
29
|
+
prf_1 = prf_tensor_1.squeeze().detach().cpu().numpy()
|
|
30
|
+
prf_2 = prf_tensor_2.squeeze().detach().cpu().numpy()
|
|
31
|
+
sigma = np.linspace(-0.025, 0.025, 51)
|
|
32
|
+
|
|
33
|
+
fig, axs = plt.subplots(1, 2, figsize=(10, 4), sharey=False)
|
|
34
|
+
|
|
35
|
+
axs[0].plot(sigma, prf_1, linestyle='-', linewidth=2)
|
|
36
|
+
axs[0].set_title(f"σ-Profile: {smiles_1}", fontsize=13)
|
|
37
|
+
axs[0].set_xlabel("σ (e/Ų)", fontsize=12)
|
|
38
|
+
axs[0].set_ylabel("P(σ)", fontsize=12)
|
|
39
|
+
axs[0].tick_params(labelsize=10)
|
|
40
|
+
|
|
41
|
+
axs[1].plot(sigma, prf_2, linestyle='-', linewidth=2, color='tab:orange')
|
|
42
|
+
axs[1].set_title(f"σ-Profile: {smiles_2}", fontsize=13)
|
|
43
|
+
axs[1].set_xlabel("σ (e/Ų)", fontsize=12)
|
|
44
|
+
axs[1].tick_params(labelsize=10)
|
|
45
|
+
|
|
46
|
+
plt.tight_layout()
|
|
47
|
+
if savepath:
|
|
48
|
+
plt.savefig(savepath, dpi=300, bbox_inches="tight")
|
|
49
|
+
plt.show()
|
|
50
|
+
|
|
51
|
+
def plot_binary_lng(x1_list, temperature, ln_gamma_1, ln_gamma_2, smiles_1, smiles_2, savepath=None):
|
|
52
|
+
plt.figure(figsize=(6, 4))
|
|
53
|
+
plt.plot(x1_list, ln_gamma_1, label='ln γ₁', linewidth=2)
|
|
54
|
+
plt.plot(x1_list, ln_gamma_2, label='ln γ₂', linewidth=2)
|
|
55
|
+
|
|
56
|
+
plt.xlabel("x₁", fontsize=12)
|
|
57
|
+
plt.ylabel("ln γ", fontsize=12)
|
|
58
|
+
plt.title(f"Binary Mixture, T = {temperature} K:\n1: {smiles_1}, 2: {smiles_2}", fontsize=13)
|
|
59
|
+
plt.xticks(fontsize=10)
|
|
60
|
+
plt.yticks(fontsize=10)
|
|
61
|
+
plt.legend(fontsize=11)
|
|
62
|
+
|
|
63
|
+
plt.xlim(0, 1)
|
|
64
|
+
y_min = min(np.min(ln_gamma_1), np.min(ln_gamma_2)) - 0.01
|
|
65
|
+
y_max = max(np.max(ln_gamma_1), np.max(ln_gamma_2)) + 0.01
|
|
66
|
+
plt.ylim(y_min, y_max)
|
|
67
|
+
|
|
68
|
+
plt.tight_layout()
|
|
69
|
+
if savepath:
|
|
70
|
+
plt.savefig(savepath, dpi=300, bbox_inches="tight")
|
|
71
|
+
plt.show()
|
|
72
|
+
|
|
73
|
+
def make_binary_df(x1_list, ln_gamma_1, ln_gamma_2, smiles_1, smiles_2, temperature):
|
|
74
|
+
df = pd.DataFrame({
|
|
75
|
+
"smiles_1": smiles_1,
|
|
76
|
+
"smiles_2": smiles_2,
|
|
77
|
+
"temperature (K)": temperature,
|
|
78
|
+
"x1": x1_list,
|
|
79
|
+
"ln_gamma_1": ln_gamma_1,
|
|
80
|
+
"ln_gamma_2": ln_gamma_2
|
|
81
|
+
})
|
|
82
|
+
return df
|