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.
Files changed (79) hide show
  1. tennetsac/__init__.py +1 -0
  2. tennetsac/ckpt_files/base.ckpt +0 -0
  3. tennetsac/ckpt_files/fine-tuned/1.ckpt +0 -0
  4. tennetsac/ckpt_files/fine-tuned/10.ckpt +0 -0
  5. tennetsac/ckpt_files/fine-tuned/2.ckpt +0 -0
  6. tennetsac/ckpt_files/fine-tuned/3.ckpt +0 -0
  7. tennetsac/ckpt_files/fine-tuned/4.ckpt +0 -0
  8. tennetsac/ckpt_files/fine-tuned/5.ckpt +0 -0
  9. tennetsac/ckpt_files/fine-tuned/6.ckpt +0 -0
  10. tennetsac/ckpt_files/fine-tuned/7.ckpt +0 -0
  11. tennetsac/ckpt_files/fine-tuned/8.ckpt +0 -0
  12. tennetsac/ckpt_files/fine-tuned/9.ckpt +0 -0
  13. tennetsac/ckpt_files/geo.ckpt +0 -0
  14. tennetsac/ckpt_files/prf.ckpt +0 -0
  15. tennetsac/core.py +88 -0
  16. tennetsac/models/Emb2Geometry.py +105 -0
  17. tennetsac/models/Emb2Profile.py +85 -0
  18. tennetsac/models/Prf2Gamma.py +120 -0
  19. tennetsac/utils/embedding.py +50 -0
  20. tennetsac/utils/model_io.py +19 -0
  21. tennetsac/utils/plotting.py +82 -0
  22. tennetsac/utils/property.py +161 -0
  23. tennetsac/utils/smi_ted_light/bert_vocab_curated.txt +2393 -0
  24. tennetsac/utils/smi_ted_light/fast_transformers/__init__.py +15 -0
  25. tennetsac/utils/smi_ted_light/fast_transformers/aggregate/__init__.py +128 -0
  26. tennetsac/utils/smi_ted_light/fast_transformers/attention/__init__.py +20 -0
  27. tennetsac/utils/smi_ted_light/fast_transformers/attention/attention_layer.py +113 -0
  28. tennetsac/utils/smi_ted_light/fast_transformers/attention/causal_linear_attention.py +116 -0
  29. tennetsac/utils/smi_ted_light/fast_transformers/attention/clustered_attention.py +195 -0
  30. tennetsac/utils/smi_ted_light/fast_transformers/attention/conditional_full_attention.py +66 -0
  31. tennetsac/utils/smi_ted_light/fast_transformers/attention/exact_topk_attention.py +88 -0
  32. tennetsac/utils/smi_ted_light/fast_transformers/attention/full_attention.py +95 -0
  33. tennetsac/utils/smi_ted_light/fast_transformers/attention/improved_clustered_attention.py +268 -0
  34. tennetsac/utils/smi_ted_light/fast_transformers/attention/improved_clustered_causal_attention.py +257 -0
  35. tennetsac/utils/smi_ted_light/fast_transformers/attention/linear_attention.py +92 -0
  36. tennetsac/utils/smi_ted_light/fast_transformers/attention/local_attention.py +101 -0
  37. tennetsac/utils/smi_ted_light/fast_transformers/attention/reformer_attention.py +166 -0
  38. tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/__init__.py +17 -0
  39. tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/registry.py +61 -0
  40. tennetsac/utils/smi_ted_light/fast_transformers/attention_registry/spec.py +126 -0
  41. tennetsac/utils/smi_ted_light/fast_transformers/builders/__init__.py +59 -0
  42. tennetsac/utils/smi_ted_light/fast_transformers/builders/attention_builders.py +139 -0
  43. tennetsac/utils/smi_ted_light/fast_transformers/builders/base.py +67 -0
  44. tennetsac/utils/smi_ted_light/fast_transformers/builders/transformer_builders.py +550 -0
  45. tennetsac/utils/smi_ted_light/fast_transformers/causal_product/__init__.py +78 -0
  46. tennetsac/utils/smi_ted_light/fast_transformers/clustering/__init__.py +0 -0
  47. tennetsac/utils/smi_ted_light/fast_transformers/clustering/hamming/__init__.py +115 -0
  48. tennetsac/utils/smi_ted_light/fast_transformers/events/__init__.py +10 -0
  49. tennetsac/utils/smi_ted_light/fast_transformers/events/event.py +51 -0
  50. tennetsac/utils/smi_ted_light/fast_transformers/events/event_dispatcher.py +92 -0
  51. tennetsac/utils/smi_ted_light/fast_transformers/events/filters.py +141 -0
  52. tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/__init__.py +12 -0
  53. tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/base.py +73 -0
  54. tennetsac/utils/smi_ted_light/fast_transformers/feature_maps/fourier_features.py +287 -0
  55. tennetsac/utils/smi_ted_light/fast_transformers/hashing/__init__.py +31 -0
  56. tennetsac/utils/smi_ted_light/fast_transformers/local_product/__init__.py +97 -0
  57. tennetsac/utils/smi_ted_light/fast_transformers/masking.py +206 -0
  58. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/__init__.py +7 -0
  59. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/_utils.py +16 -0
  60. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/__init__.py +16 -0
  61. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/__init__.py +30 -0
  62. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/attention_layer.py +105 -0
  63. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/full_attention.py +75 -0
  64. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/cross_attention/linear_attention.py +79 -0
  65. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/__init__.py +30 -0
  66. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/attention_layer.py +96 -0
  67. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/full_attention.py +83 -0
  68. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/attention/self_attention/linear_attention.py +110 -0
  69. tennetsac/utils/smi_ted_light/fast_transformers/recurrent/transformers.py +279 -0
  70. tennetsac/utils/smi_ted_light/fast_transformers/sparse_product/__init__.py +399 -0
  71. tennetsac/utils/smi_ted_light/fast_transformers/transformers.py +294 -0
  72. tennetsac/utils/smi_ted_light/fast_transformers/utils.py +33 -0
  73. tennetsac/utils/smi_ted_light/fast_transformers/weight_mapper.py +273 -0
  74. tennetsac/utils/smi_ted_light/load.py +680 -0
  75. tennetsac/utils/smiles.py +28 -0
  76. tennetsac-0.1.1.dist-info/METADATA +55 -0
  77. tennetsac-0.1.1.dist-info/RECORD +79 -0
  78. tennetsac-0.1.1.dist-info/WHEEL +5 -0
  79. 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