heiwa 0.1.0__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.
- heiwa-0.1.0/PKG-INFO +13 -0
- heiwa-0.1.0/pyproject.toml +40 -0
- heiwa-0.1.0/setup.cfg +4 -0
- heiwa-0.1.0/src/heiwa/__init__.py +43 -0
- heiwa-0.1.0/src/heiwa/batch_run.py +67 -0
- heiwa-0.1.0/src/heiwa/generate/esm3_gen.py +159 -0
- heiwa-0.1.0/src/heiwa/generate/generate.py +105 -0
- heiwa-0.1.0/src/heiwa/generate/generate_utils.py +18 -0
- heiwa-0.1.0/src/heiwa/generate/modeller_script.py +89 -0
- heiwa-0.1.0/src/heiwa/hardware.py +24 -0
- heiwa-0.1.0/src/heiwa/loaders/dataloader.py +71 -0
- heiwa-0.1.0/src/heiwa/loaders/model_loader.py +31 -0
- heiwa-0.1.0/src/heiwa/loaders/tensor_loader.py +91 -0
- heiwa-0.1.0/src/heiwa/model_assets/ExposedResiduePrediction_epoch_final_model.pth.gz +0 -0
- heiwa-0.1.0/src/heiwa/model_assets/HEIWA_Hybridizer_epoch_final_model.pth.gz +0 -0
- heiwa-0.1.0/src/heiwa/model_assets/HEIWA_Imp_epoch_final_model.pth.gz +0 -0
- heiwa-0.1.0/src/heiwa/model_assets/TIM_epoch_final_model.pth.gz +0 -0
- heiwa-0.1.0/src/heiwa/model_assets/Thermo_epoch_final_model.pth.gz +0 -0
- heiwa-0.1.0/src/heiwa/model_assets/avg1536_wt.pt +0 -0
- heiwa-0.1.0/src/heiwa/model_assets/eigvec.pt +0 -0
- heiwa-0.1.0/src/heiwa/model_assets/logit_ranks.npy +0 -0
- heiwa-0.1.0/src/heiwa/model_assets/rmsf_ranks.npy +0 -0
- heiwa-0.1.0/src/heiwa/model_assets/sd1536_wt.pt +0 -0
- heiwa-0.1.0/src/heiwa/model_run/batch_accumulator.py +58 -0
- heiwa-0.1.0/src/heiwa/model_run/epoch_run.py +117 -0
- heiwa-0.1.0/src/heiwa/model_run/rmsf_convert.py +11 -0
- heiwa-0.1.0/src/heiwa/models/Water3P_Implicit.py +54 -0
- heiwa-0.1.0/src/heiwa/models/basic_layers.py +89 -0
- heiwa-0.1.0/src/heiwa/models/data_containers.py +8 -0
- heiwa-0.1.0/src/heiwa/models/file_handling.py +27 -0
- heiwa-0.1.0/src/heiwa/models/model.py +77 -0
- heiwa-0.1.0/src/heiwa/models/orthogonalizer.py +20 -0
- heiwa-0.1.0/src/heiwa/models/part_HEIWA_Hybridizer.py +41 -0
- heiwa-0.1.0/src/heiwa/models/part_HEIWA_Imp.py +139 -0
- heiwa-0.1.0/src/heiwa/models/part_SASA.py +34 -0
- heiwa-0.1.0/src/heiwa/models/part_TIM.py +52 -0
- heiwa-0.1.0/src/heiwa/models/part_thermo.py +44 -0
- heiwa-0.1.0/src/heiwa/predict.py +100 -0
- heiwa-0.1.0/src/heiwa/readers/read_cif.py +90 -0
- heiwa-0.1.0/src/heiwa/readers/read_dssp.py +76 -0
- heiwa-0.1.0/src/heiwa/readers/read_pdb.py +119 -0
- heiwa-0.1.0/src/heiwa/readers/read_utils.py +18 -0
- heiwa-0.1.0/src/heiwa/readers/reader.py +32 -0
- heiwa-0.1.0/src/heiwa/replace_esm3/blocks.py +164 -0
- heiwa-0.1.0/src/heiwa/replace_esm3/encoding.py +246 -0
- heiwa-0.1.0/src/heiwa/replace_esm3/esm3.py +652 -0
- heiwa-0.1.0/src/heiwa/replace_esm3/generation.py +845 -0
- heiwa-0.1.0/src/heiwa/replace_esm3/sequence_tokenizer.py +137 -0
- heiwa-0.1.0/src/heiwa/replace_esm3/transformer_stack.py +105 -0
- heiwa-0.1.0/src/heiwa/utils.py +42 -0
- heiwa-0.1.0/src/heiwa.egg-info/PKG-INFO +13 -0
- heiwa-0.1.0/src/heiwa.egg-info/SOURCES.txt +53 -0
- heiwa-0.1.0/src/heiwa.egg-info/dependency_links.txt +1 -0
- heiwa-0.1.0/src/heiwa.egg-info/requires.txt +10 -0
- heiwa-0.1.0/src/heiwa.egg-info/top_level.txt +1 -0
heiwa-0.1.0/PKG-INFO
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: heiwa
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Requires-Dist: biotite
|
|
5
|
+
Requires-Dist: biopython
|
|
6
|
+
Requires-Dist: numpy
|
|
7
|
+
Requires-Dist: pandas
|
|
8
|
+
Requires-Dist: torch>=2.5.1
|
|
9
|
+
Requires-Dist: tqdm
|
|
10
|
+
Requires-Dist: esm==3.2.1
|
|
11
|
+
Requires-Dist: PyCifRW
|
|
12
|
+
Requires-Dist: transformers
|
|
13
|
+
Requires-Dist: tokenizers
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
requires = ["setuptools>=75"]
|
|
3
|
+
build-backend = "setuptools.build_meta"
|
|
4
|
+
|
|
5
|
+
[project]
|
|
6
|
+
name = "heiwa"
|
|
7
|
+
version = "0.1.0"
|
|
8
|
+
|
|
9
|
+
dependencies = [
|
|
10
|
+
"biotite",
|
|
11
|
+
"biopython",
|
|
12
|
+
"numpy",
|
|
13
|
+
"pandas",
|
|
14
|
+
"torch>=2.5.1",
|
|
15
|
+
"tqdm",
|
|
16
|
+
"esm==3.2.1",
|
|
17
|
+
"PyCifRW",
|
|
18
|
+
"transformers",
|
|
19
|
+
"tokenizers"
|
|
20
|
+
]
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
[tool.setuptools.packages.find]
|
|
24
|
+
where = ["src"]
|
|
25
|
+
|
|
26
|
+
[tool.setuptools.package-data]
|
|
27
|
+
heiwa = ["model_assets/*"]
|
|
28
|
+
|
|
29
|
+
# requires-python = ">=3.11"
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
# - python
|
|
33
|
+
# - numpy
|
|
34
|
+
# - pandas
|
|
35
|
+
# - pytorch>=2.5.1
|
|
36
|
+
# - pytorch-cuda>=11.8
|
|
37
|
+
# - tqdm
|
|
38
|
+
# - pycifrw
|
|
39
|
+
# - transformers==4.46.3
|
|
40
|
+
# - tokenizers==0.20
|
heiwa-0.1.0/setup.cfg
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
# Copyright (c) 2026 Ernest Ngai Hei Ho
|
|
2
|
+
# Email: hoernesta@gmail.com
|
|
3
|
+
# Licensed under the MIT License.
|
|
4
|
+
__version__="0.1.0"
|
|
5
|
+
from .loaders.model_loader import load_solv_encoder, load_thermo
|
|
6
|
+
from .hardware import Hardware_Setting
|
|
7
|
+
# from .esm_gen import gen_embed
|
|
8
|
+
from .predict import predict
|
|
9
|
+
from .generate.generate import de_novo, gen_embed
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
## Replace esm3 script during import
|
|
13
|
+
import shutil
|
|
14
|
+
from importlib.resources import files
|
|
15
|
+
|
|
16
|
+
_source_folder="replace_esm3"
|
|
17
|
+
_package_path = files("heiwa")
|
|
18
|
+
shutil.copy(
|
|
19
|
+
_package_path.joinpath(_source_folder, "esm3.py"),
|
|
20
|
+
_package_path.parent.joinpath("esm/models/esm3.py")
|
|
21
|
+
)
|
|
22
|
+
shutil.copy(
|
|
23
|
+
_package_path.joinpath(_source_folder, "generation.py"),
|
|
24
|
+
_package_path.parent.joinpath("esm/utils/generation.py")
|
|
25
|
+
)
|
|
26
|
+
# shutil.copy(
|
|
27
|
+
# _package_path.joinpath(_source_folder, "sequence_tokenizer.py"),
|
|
28
|
+
# _package_path.parent.joinpath("esm/tokenization/sequence_tokenizer.py")
|
|
29
|
+
# )
|
|
30
|
+
# shutil.copy(
|
|
31
|
+
# _package_path.joinpath(_source_folder, "encoding.py"),
|
|
32
|
+
# _package_path.parent.joinpath("esm/utils/encoding.py")
|
|
33
|
+
# )
|
|
34
|
+
shutil.copy(
|
|
35
|
+
_package_path.joinpath(_source_folder, "transformer_stack.py"),
|
|
36
|
+
_package_path.parent.joinpath("esm/layers/transformer_stack.py")
|
|
37
|
+
)
|
|
38
|
+
shutil.copy(
|
|
39
|
+
_package_path.joinpath(_source_folder, "blocks.py"),
|
|
40
|
+
_package_path.parent.joinpath("esm/layers/blocks.py")
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
del _source_folder, _package_path
|
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
from heiwa.loaders.tensor_loader import concat_tensors_and_pad
|
|
3
|
+
from heiwa.loaders.dataloader import compile_data_single
|
|
4
|
+
from heiwa.hardware import Hardware_Setting
|
|
5
|
+
from heiwa.model_run.epoch_run import EpochRunner
|
|
6
|
+
from heiwa.utils import ensure_3d_tensor
|
|
7
|
+
|
|
8
|
+
def batch_run(
|
|
9
|
+
x:torch.Tensor|list[torch.Tensor],
|
|
10
|
+
hw_setting:Hardware_Setting|None=None,
|
|
11
|
+
heiwa_model=None,
|
|
12
|
+
heiwa_orthogonalizer=None,
|
|
13
|
+
thermo_model=None,
|
|
14
|
+
out_folder=".",
|
|
15
|
+
):
|
|
16
|
+
# assert target.lower() in ("rmsf", "thermo")
|
|
17
|
+
|
|
18
|
+
# ========= Prepare dataset here =========
|
|
19
|
+
# m = None
|
|
20
|
+
# l = None
|
|
21
|
+
if isinstance(x, list):
|
|
22
|
+
print("hi")
|
|
23
|
+
dataset = concat_tensors_and_pad(
|
|
24
|
+
embeds=x,
|
|
25
|
+
# create_l=target.lower()=="thermo",
|
|
26
|
+
hw_setting=hw_setting,
|
|
27
|
+
)
|
|
28
|
+
# print("x_unorthogonalized_list", len(x), x)
|
|
29
|
+
# print("x_unorthogonalized", dataset["X"].size(), dataset["X"])
|
|
30
|
+
# print("l", l.size(), l)
|
|
31
|
+
else:
|
|
32
|
+
print("hi single x")
|
|
33
|
+
x = ensure_3d_tensor(x=x)
|
|
34
|
+
# if target.lower()=="thermo":
|
|
35
|
+
# print(f"Task is definitely thermo")
|
|
36
|
+
l = torch.tensor(
|
|
37
|
+
[x.size(1)],
|
|
38
|
+
dtype=torch.float32,
|
|
39
|
+
device='cpu'
|
|
40
|
+
)
|
|
41
|
+
dataset={"X":x, "L":l}
|
|
42
|
+
print("x_unorthogonalized", x.size(), x)
|
|
43
|
+
print("dataset", dataset.keys())
|
|
44
|
+
|
|
45
|
+
# ========= Prepare models, hardware, epoch_runner =========
|
|
46
|
+
assert heiwa_model is not None and thermo_model is not None
|
|
47
|
+
|
|
48
|
+
hw_setting.switch_on_gpu_rs()
|
|
49
|
+
|
|
50
|
+
epoch_runner = EpochRunner(out_folder=out_folder)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
data_loader = compile_data_single(
|
|
54
|
+
hw_setting=hw_setting,
|
|
55
|
+
dataset=dataset,
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
# ========= Run ==========================
|
|
59
|
+
print(f"Start running batches")
|
|
60
|
+
|
|
61
|
+
# print("dataset", data_loader.dataset)
|
|
62
|
+
epoch_runner.run(
|
|
63
|
+
data_source=data_loader,
|
|
64
|
+
heiwa_model=heiwa_model,
|
|
65
|
+
heiwa_orthogonalizer=heiwa_orthogonalizer,
|
|
66
|
+
thermo_model=thermo_model,
|
|
67
|
+
)
|
|
@@ -0,0 +1,159 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import torch
|
|
3
|
+
import torch.nn as nn
|
|
4
|
+
import pickle
|
|
5
|
+
import pandas as pd
|
|
6
|
+
# import traceback
|
|
7
|
+
from heiwa.readers.reader import read_file
|
|
8
|
+
from heiwa.generate.generate_utils import convert_idx_to_prompt
|
|
9
|
+
from esm.sdk.api import GenerationConfig, ESMProtein
|
|
10
|
+
# from heiwa.generate.modeller_script import fill_atom_to_pdb_in_single
|
|
11
|
+
|
|
12
|
+
tracks = [
|
|
13
|
+
{'structure':16},
|
|
14
|
+
{'secondary_structure':8},
|
|
15
|
+
{'sasa':4},
|
|
16
|
+
{'structure':8},
|
|
17
|
+
{'sequence':8}
|
|
18
|
+
]
|
|
19
|
+
|
|
20
|
+
def esm3_generation(
|
|
21
|
+
real_tracks, protein,
|
|
22
|
+
esm3_model, heiwa_model, heiwa_orthogonalizer,
|
|
23
|
+
seq_prompt, crit_resi,
|
|
24
|
+
out_filename
|
|
25
|
+
):
|
|
26
|
+
if protein == "max_len_exceeded":
|
|
27
|
+
pass
|
|
28
|
+
fname = None
|
|
29
|
+
for i, track in enumerate(real_tracks):
|
|
30
|
+
# print(track)
|
|
31
|
+
if i == len(real_tracks)-1:
|
|
32
|
+
fname = out_filename
|
|
33
|
+
|
|
34
|
+
track_k = list(track.keys())[0]
|
|
35
|
+
track_v = list(track.values())[0]
|
|
36
|
+
# print(track)
|
|
37
|
+
# print(fname)
|
|
38
|
+
condition_on_coordinates_only = True if track_k == 'structure' else False
|
|
39
|
+
|
|
40
|
+
if track_k=='sequence':
|
|
41
|
+
protein.sequence = seq_prompt
|
|
42
|
+
# print(seq_prompt)
|
|
43
|
+
elif track_k=='structure' and track_v == 8:
|
|
44
|
+
protein = esm3_model.generate(protein, GenerationConfig(track=track_k, num_steps=track_v, \
|
|
45
|
+
temperature=0.7,
|
|
46
|
+
temperature_annealing=True, \
|
|
47
|
+
condition_on_coordinates_only = condition_on_coordinates_only), \
|
|
48
|
+
filename=fname,
|
|
49
|
+
layer_i=-1,
|
|
50
|
+
heiwa_model=heiwa_model,
|
|
51
|
+
heiwa_normalizer=heiwa_orthogonalizer)
|
|
52
|
+
continue
|
|
53
|
+
# print(len(protein.sequence), len(protein.coordinates), protein.sequence)
|
|
54
|
+
protein = esm3_model.generate(protein, GenerationConfig(track=track_k, num_steps=track_v, \
|
|
55
|
+
temperature=0.7,
|
|
56
|
+
temperature_annealing=True, \
|
|
57
|
+
condition_on_coordinates_only = condition_on_coordinates_only), \
|
|
58
|
+
filename=fname,
|
|
59
|
+
layer_i=-1)
|
|
60
|
+
return protein
|
|
61
|
+
|
|
62
|
+
def check_generation_success(outPath:str|list[str]):
|
|
63
|
+
"""Check where pdb and pkl exists at the right path"""
|
|
64
|
+
if isinstance(outPath, str):
|
|
65
|
+
with open(f"./generation_report.txt", "a+") as file:
|
|
66
|
+
if os.path.exists(outPath):
|
|
67
|
+
file.write(f"Generation of {outPath} successful.\n")
|
|
68
|
+
else:
|
|
69
|
+
file.write(f"Generation of {outPath} failed.\n")
|
|
70
|
+
elif isinstance(outPath, list):
|
|
71
|
+
with open(f"./generation_report.txt", "a+") as file:
|
|
72
|
+
for path in outPath:
|
|
73
|
+
if os.path.exists(outPath):
|
|
74
|
+
file.write(f"Generation of {path} successful.\n")
|
|
75
|
+
else:
|
|
76
|
+
file.write(f"Generation of {path} failed.\n")
|
|
77
|
+
|
|
78
|
+
def save_protein(protein, folder_out, out_file_name):
|
|
79
|
+
with open(f"{folder_out}/{out_file_name}.pkl", 'wb') as file:
|
|
80
|
+
pickle.dump(protein, file)
|
|
81
|
+
protein.to_pdb(f"{folder_out}/{out_file_name}.pdb")
|
|
82
|
+
|
|
83
|
+
def generate_proteins(
|
|
84
|
+
task,
|
|
85
|
+
hw_setting,
|
|
86
|
+
esm3_model,
|
|
87
|
+
heiwa_model=None, heiwa_orthogonalizer=None, seq_prompt=None, crit_resi=None, idx_style:int|None=None,
|
|
88
|
+
inFile:str|None=None,
|
|
89
|
+
protein:ESMProtein|None=None,
|
|
90
|
+
cleave:int=0,
|
|
91
|
+
out_file_name:str|None=None,
|
|
92
|
+
folder_out=".", \
|
|
93
|
+
debug_path: str="",
|
|
94
|
+
):
|
|
95
|
+
assert task in ("gen_embed", "gen_de_novo")
|
|
96
|
+
assert esm3_model
|
|
97
|
+
assert out_file_name
|
|
98
|
+
stream = torch.cuda.Stream(device="cuda")
|
|
99
|
+
if heiwa_model:
|
|
100
|
+
assert heiwa_orthogonalizer
|
|
101
|
+
|
|
102
|
+
if task == "gen_embed":
|
|
103
|
+
out_file_type="pt"
|
|
104
|
+
real_tracks = tracks[:-1]
|
|
105
|
+
else:
|
|
106
|
+
assert heiwa_model and heiwa_orthogonalizer
|
|
107
|
+
assert seq_prompt is not None or (crit_resi is not None and idx_style is not None)
|
|
108
|
+
out_file_type="pdb"
|
|
109
|
+
real_tracks = tracks
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
if inFile:
|
|
113
|
+
protein = read_file(
|
|
114
|
+
filename=inFile,
|
|
115
|
+
cleave=cleave,
|
|
116
|
+
debug_path=debug_path,
|
|
117
|
+
)
|
|
118
|
+
# print("b", type(protein))
|
|
119
|
+
|
|
120
|
+
if task=="gen_de_novo" and seq_prompt is None:
|
|
121
|
+
seq_prompt = convert_idx_to_prompt(
|
|
122
|
+
idx=crit_resi,
|
|
123
|
+
cleave=cleave,
|
|
124
|
+
seq_wt=protein.sequence,
|
|
125
|
+
length=len(protein.sequence),
|
|
126
|
+
idx_style=idx_style,
|
|
127
|
+
)
|
|
128
|
+
# print(seq_prompt)
|
|
129
|
+
|
|
130
|
+
hw_setting.switch_on_gpu_rs()
|
|
131
|
+
|
|
132
|
+
if out_file_name.count('.') > 0:
|
|
133
|
+
out_file_name = out_file_name[:out_file_name.rfind('.')]
|
|
134
|
+
if not os.path.exists(f"{folder_out}/{out_file_name}.{out_file_type}"):
|
|
135
|
+
protein = esm3_generation(
|
|
136
|
+
real_tracks=real_tracks,
|
|
137
|
+
protein=protein,
|
|
138
|
+
seq_prompt=seq_prompt,
|
|
139
|
+
crit_resi=crit_resi,
|
|
140
|
+
esm3_model=esm3_model, heiwa_model=heiwa_model, heiwa_orthogonalizer=heiwa_orthogonalizer,
|
|
141
|
+
out_filename=f"{folder_out}/{out_file_name}"
|
|
142
|
+
)
|
|
143
|
+
|
|
144
|
+
if task=="gen_de_novo":
|
|
145
|
+
save_protein(protein=protein,
|
|
146
|
+
folder_out=folder_out,
|
|
147
|
+
out_file_name=out_file_name)
|
|
148
|
+
|
|
149
|
+
# fill_atom_to_pdb_in_single(
|
|
150
|
+
# base_name=file_name,
|
|
151
|
+
# folder_out=folder_out,
|
|
152
|
+
# sequence=protein.sequence,
|
|
153
|
+
# length=len(protein.sequence)
|
|
154
|
+
# )
|
|
155
|
+
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
|
|
@@ -0,0 +1,105 @@
|
|
|
1
|
+
import torch.nn as nn
|
|
2
|
+
from esm.sdk.api import ESMProtein
|
|
3
|
+
from heiwa.generate.esm3_gen import generate_proteins
|
|
4
|
+
from heiwa.utils import ensure_dirs
|
|
5
|
+
|
|
6
|
+
def gen_embed(
|
|
7
|
+
esm3_model,
|
|
8
|
+
hw_setting,
|
|
9
|
+
inFile:str|None=None,
|
|
10
|
+
protein:ESMProtein|None=None,
|
|
11
|
+
cleave:int=0,
|
|
12
|
+
outPath:str="./no_name.pt",
|
|
13
|
+
debug_path:str="",
|
|
14
|
+
):
|
|
15
|
+
"""1 protein at once"""
|
|
16
|
+
if outPath.count('/') > 0:
|
|
17
|
+
sep = outPath.rfind('/')
|
|
18
|
+
folder_out = outPath[:sep]
|
|
19
|
+
outPath = outPath[sep+1:]
|
|
20
|
+
else:
|
|
21
|
+
folder_out = "."
|
|
22
|
+
ensure_dirs(folder_out)
|
|
23
|
+
|
|
24
|
+
if inFile:
|
|
25
|
+
generate_proteins(
|
|
26
|
+
task="gen_embed",
|
|
27
|
+
hw_setting=hw_setting,
|
|
28
|
+
inFile=inFile,
|
|
29
|
+
out_file_name=outPath,
|
|
30
|
+
cleave=cleave,
|
|
31
|
+
esm3_model=esm3_model,
|
|
32
|
+
folder_out=folder_out,
|
|
33
|
+
debug_path=debug_path,
|
|
34
|
+
)
|
|
35
|
+
else:
|
|
36
|
+
assert protein
|
|
37
|
+
generate_proteins(
|
|
38
|
+
task="gen_embed",
|
|
39
|
+
hw_setting=hw_setting,
|
|
40
|
+
protein=protein,
|
|
41
|
+
cleave=cleave,
|
|
42
|
+
out_file_name=outPath,
|
|
43
|
+
esm3_model=esm3_model,
|
|
44
|
+
folder_out=folder_out,
|
|
45
|
+
debug_path=debug_path,
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
def de_novo(
|
|
49
|
+
esm3_model,
|
|
50
|
+
hw_setting,
|
|
51
|
+
heiwa_model:nn.Module|None=None,
|
|
52
|
+
heiwa_orthogonalizer:nn.Module|None=None,
|
|
53
|
+
inFile:str|None=None,
|
|
54
|
+
seq_prompt:str|None=None,
|
|
55
|
+
crit_resi:list[int]|None=None,
|
|
56
|
+
idx_style:int|None=None,
|
|
57
|
+
cleave:int=0,
|
|
58
|
+
protein=None,
|
|
59
|
+
outPath:str="./no_name.pdb",
|
|
60
|
+
debug_path:str="",
|
|
61
|
+
):
|
|
62
|
+
"""1 protein at once"""
|
|
63
|
+
assert not (seq_prompt is not None and (crit_resi is not None or idx_style))
|
|
64
|
+
|
|
65
|
+
if outPath.count('/') > 0:
|
|
66
|
+
sep = outPath.rfind('/')
|
|
67
|
+
folder_out = outPath[:sep]
|
|
68
|
+
outPath = outPath[sep+1:]
|
|
69
|
+
else:
|
|
70
|
+
folder_out = "."
|
|
71
|
+
ensure_dirs(folder_out)
|
|
72
|
+
|
|
73
|
+
if inFile:
|
|
74
|
+
generate_proteins(
|
|
75
|
+
task="gen_de_novo",
|
|
76
|
+
hw_setting=hw_setting,
|
|
77
|
+
inFile=inFile,
|
|
78
|
+
cleave=cleave,
|
|
79
|
+
seq_prompt=seq_prompt,
|
|
80
|
+
crit_resi=crit_resi,
|
|
81
|
+
idx_style=idx_style,
|
|
82
|
+
out_file_name=outPath,
|
|
83
|
+
esm3_model=esm3_model,
|
|
84
|
+
heiwa_model=heiwa_model,
|
|
85
|
+
heiwa_orthogonalizer=heiwa_orthogonalizer,
|
|
86
|
+
folder_out=folder_out,
|
|
87
|
+
debug_path=debug_path,
|
|
88
|
+
)
|
|
89
|
+
else:
|
|
90
|
+
assert protein
|
|
91
|
+
generate_proteins(
|
|
92
|
+
task="gen_de_novo",
|
|
93
|
+
hw_setting=hw_setting,
|
|
94
|
+
protein=protein,
|
|
95
|
+
cleave=cleave,
|
|
96
|
+
seq_prompt=seq_prompt,
|
|
97
|
+
crit_resi=crit_resi,
|
|
98
|
+
idx_style=idx_style,
|
|
99
|
+
out_file_name=outPath,
|
|
100
|
+
esm3_model=esm3_model,
|
|
101
|
+
heiwa_model=heiwa_model,
|
|
102
|
+
heiwa_orthogonalizer=heiwa_orthogonalizer,
|
|
103
|
+
folder_out=folder_out,
|
|
104
|
+
debug_path=debug_path,
|
|
105
|
+
)
|
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
|
|
2
|
+
def convert_idx_to_prompt(
|
|
3
|
+
idx:list[int]|tuple[int], # only this is affected by signal peptide
|
|
4
|
+
seq_wt:str,
|
|
5
|
+
length:int,
|
|
6
|
+
cleave:int=0,
|
|
7
|
+
idx_style=0,
|
|
8
|
+
):
|
|
9
|
+
seq_prompt = ""
|
|
10
|
+
i_last=0
|
|
11
|
+
# print(seq_wt)
|
|
12
|
+
idx = sorted([i-cleave for i in idx])
|
|
13
|
+
for i in idx:
|
|
14
|
+
# print(i-i_last-idx_style, seq_wt[i-idx_style])
|
|
15
|
+
seq_prompt += "_"*(i-i_last-idx_style) + seq_wt[i-idx_style]
|
|
16
|
+
i_last = i
|
|
17
|
+
seq_prompt += "_"*(length - i_last + 1 - idx_style)
|
|
18
|
+
return seq_prompt
|
|
@@ -0,0 +1,89 @@
|
|
|
1
|
+
from modeller import *
|
|
2
|
+
from modeller.automodel import *
|
|
3
|
+
from modeller.parallel import *
|
|
4
|
+
import os
|
|
5
|
+
|
|
6
|
+
env = Environ()
|
|
7
|
+
default_output = "target.B99990001.pdb"
|
|
8
|
+
|
|
9
|
+
def alignment_writer(sequence, length, base_name, folder_out):
|
|
10
|
+
"""Write alignment file for MODELLER
|
|
11
|
+
Args:
|
|
12
|
+
row: pandas DataFrame row containing 'Sequence_final', 'Filename', 'Y_thermo'
|
|
13
|
+
"""
|
|
14
|
+
file_out_name = f"{folder_out}/{base_name}.ali"
|
|
15
|
+
|
|
16
|
+
# Prepare the alignment content
|
|
17
|
+
alignment_content = f""">P1;backbone
|
|
18
|
+
structureX:{base_name}.pdb:1:A:{length}:A::::
|
|
19
|
+
{'-' * length}*
|
|
20
|
+
|
|
21
|
+
>P1;target
|
|
22
|
+
sequence:target::::::::
|
|
23
|
+
{sequence}*
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
# Write the alignment file
|
|
27
|
+
with open(file_out_name, 'w') as f:
|
|
28
|
+
f.write(alignment_content)
|
|
29
|
+
|
|
30
|
+
print(f"Created alignment file: {file_out_name}")
|
|
31
|
+
return file_out_name
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def fill_atom_to_pdb(base_name, folder_out):
|
|
35
|
+
"""Build full-atom model from backbone-only PDB
|
|
36
|
+
Args:
|
|
37
|
+
row: pandas DataFrame row containing 'Filename', 'Y_thermo'
|
|
38
|
+
"""
|
|
39
|
+
pdb_path = f"{folder_out}/{base_name}.pdb"
|
|
40
|
+
ali_file = f"{folder_out}/{base_name}.ali"
|
|
41
|
+
ali_final = f"{folder_out}/{base_name}_final.ali"
|
|
42
|
+
|
|
43
|
+
# Read the template structure
|
|
44
|
+
aln = Alignment(env)
|
|
45
|
+
mdl = Model(env, file=pdb_path.replace('.pdb', ''), model_segment=('FIRST:A', 'LAST:A'))
|
|
46
|
+
aln.append_model(mdl, align_codes='backbone', atom_files=pdb_path)
|
|
47
|
+
aln.append(file=ali_file, align_codes='target')
|
|
48
|
+
aln.align2d()
|
|
49
|
+
aln.write(file=ali_final, alignment_format='PIR')
|
|
50
|
+
|
|
51
|
+
# Build the model with side chains
|
|
52
|
+
a = AutoModel(env,
|
|
53
|
+
alnfile=ali_final,
|
|
54
|
+
knowns='backbone',
|
|
55
|
+
sequence='target')
|
|
56
|
+
a.starting_model = 1
|
|
57
|
+
a.ending_model = 1
|
|
58
|
+
|
|
59
|
+
# Set output directory
|
|
60
|
+
a.outputs = ['MODELS'] # Only output final model
|
|
61
|
+
|
|
62
|
+
a.make()
|
|
63
|
+
|
|
64
|
+
# Rename output to something more meaningful
|
|
65
|
+
# MODELLER creates: target.B99990001.pdb
|
|
66
|
+
final_output = f"{folder_out}/{base_name}_full.pdb"
|
|
67
|
+
|
|
68
|
+
if os.path.exists(default_output):
|
|
69
|
+
os.rename(default_output, final_output)
|
|
70
|
+
print(f"Created full-atom model: {final_output}")
|
|
71
|
+
for path in [pdb_path, ali_file, ali_final]:
|
|
72
|
+
os.remove(path)
|
|
73
|
+
else:
|
|
74
|
+
print(f"Warning: Expected output {default_output} not found!")
|
|
75
|
+
|
|
76
|
+
return final_output
|
|
77
|
+
|
|
78
|
+
def fill_atom_to_pdb_in_single(base_name, folder_out, sequence, length):
|
|
79
|
+
alignment_writer(
|
|
80
|
+
sequence=sequence, length=length, base_name=base_name, folder_out=folder_out
|
|
81
|
+
)
|
|
82
|
+
print("wrote alignment")
|
|
83
|
+
|
|
84
|
+
fill_atom_to_pdb(
|
|
85
|
+
base_name, folder_out
|
|
86
|
+
)
|
|
87
|
+
print("pdb filled")
|
|
88
|
+
|
|
89
|
+
|
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
import torch
|
|
3
|
+
# import numpy as np
|
|
4
|
+
|
|
5
|
+
@dataclass
|
|
6
|
+
class Hardware_Setting:
|
|
7
|
+
cpu: int = 1
|
|
8
|
+
batch_size : int = 2
|
|
9
|
+
gpu_rs: int|None = None
|
|
10
|
+
# cpu_rs: int
|
|
11
|
+
|
|
12
|
+
def __str__(self):
|
|
13
|
+
cout = "HEIWA hardware setting:\n"
|
|
14
|
+
cout += f"cpu = {self.cpu}\n"
|
|
15
|
+
cout += f"batch_size = {self.batch_size}\n"
|
|
16
|
+
cout += f"gpu_rs = {self.gpu_rs}\n"
|
|
17
|
+
# cout += f"cpu_rs = {self.cpu_rs}\n"
|
|
18
|
+
return cout
|
|
19
|
+
|
|
20
|
+
def switch_on_gpu_rs(self):
|
|
21
|
+
if self.gpu_rs is not None:
|
|
22
|
+
torch.manual_seed(self.gpu_rs)
|
|
23
|
+
|
|
24
|
+
|
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import numpy as np
|
|
3
|
+
from torch.utils.data import Dataset, DataLoader
|
|
4
|
+
import numpy as np
|
|
5
|
+
from typing import Tuple, Dict
|
|
6
|
+
from heiwa.hardware import Hardware_Setting
|
|
7
|
+
|
|
8
|
+
class MyDataset(Dataset):
|
|
9
|
+
def __init__(self, prefix):
|
|
10
|
+
""" The real initlization is called by cls._from_dict """
|
|
11
|
+
self.DATAFILE_PREFIX = tuple(prefix)
|
|
12
|
+
print("Tensor prefix", self.DATAFILE_PREFIX)
|
|
13
|
+
|
|
14
|
+
@classmethod
|
|
15
|
+
def _from_dict(cls, all_inputs:dict):
|
|
16
|
+
""" This triggers __init__ """
|
|
17
|
+
dataset = MyDataset(prefix=all_inputs.keys())
|
|
18
|
+
for key, value in all_inputs.items():
|
|
19
|
+
setattr(dataset, key, value)
|
|
20
|
+
if key == "X":
|
|
21
|
+
print(f"{key}: {len(value)}")
|
|
22
|
+
else:
|
|
23
|
+
print(f"{key}: {value}")
|
|
24
|
+
return dataset
|
|
25
|
+
|
|
26
|
+
def __str__(self):
|
|
27
|
+
cout = ""
|
|
28
|
+
for k in "XLM":
|
|
29
|
+
cout += f"{k}:{getattr(self, k, None)}\n"
|
|
30
|
+
return cout
|
|
31
|
+
|
|
32
|
+
def __len__(self) -> int:
|
|
33
|
+
if hasattr(self, "X"):
|
|
34
|
+
return len(self.X)
|
|
35
|
+
|
|
36
|
+
def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
|
|
37
|
+
output = {}
|
|
38
|
+
output['model_input'] = {k: getattr(self, k)[idx] for k in self.DATAFILE_PREFIX}
|
|
39
|
+
return output
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def compile_data_single(
|
|
44
|
+
hw_setting:Hardware_Setting,
|
|
45
|
+
dataset:dict
|
|
46
|
+
) -> DataLoader:
|
|
47
|
+
""" Create DataLoader objects for train, validation, and test sets.
|
|
48
|
+
|
|
49
|
+
Args:
|
|
50
|
+
target: 'test', 'inference', 'train' if fisher
|
|
51
|
+
dataset: dictionary of individual raw tensors
|
|
52
|
+
params: all Parameters
|
|
53
|
+
|
|
54
|
+
Returns:
|
|
55
|
+
output_dataset """
|
|
56
|
+
# Create datasets
|
|
57
|
+
# print(dataset)
|
|
58
|
+
output_dataset = MyDataset._from_dict(
|
|
59
|
+
dataset # {k:dataset[k] for k in dataset}
|
|
60
|
+
) #
|
|
61
|
+
|
|
62
|
+
output_dataset = DataLoader(
|
|
63
|
+
output_dataset,
|
|
64
|
+
batch_size=hw_setting.batch_size,
|
|
65
|
+
shuffle=False,
|
|
66
|
+
num_workers=hw_setting.cpu,
|
|
67
|
+
pin_memory=True
|
|
68
|
+
)
|
|
69
|
+
print(f"Dataloader for compiled")
|
|
70
|
+
return output_dataset
|
|
71
|
+
|
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
from heiwa.models.model import HEIWA3H
|
|
3
|
+
from heiwa.models.orthogonalizer import HEIWA_Orthogonalizer
|
|
4
|
+
from heiwa.models.part_thermo import ThermoPred
|
|
5
|
+
from heiwa.utils import asset_path
|
|
6
|
+
|
|
7
|
+
def load_esm3():
|
|
8
|
+
with torch.no_grad():
|
|
9
|
+
return torch.load(
|
|
10
|
+
asset_path.joinpath("esm3.pth"),
|
|
11
|
+
map_location="cuda",
|
|
12
|
+
weights_only=False
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
def load_solv_encoder(hw_setting):
|
|
16
|
+
with torch.no_grad():
|
|
17
|
+
heiwa = HEIWA3H(hw_setting=hw_setting)
|
|
18
|
+
|
|
19
|
+
# heiwa = torch.load(
|
|
20
|
+
# path,
|
|
21
|
+
# map_location="cuda",
|
|
22
|
+
# weights_only=False
|
|
23
|
+
# )
|
|
24
|
+
|
|
25
|
+
heiwa.to("cuda")
|
|
26
|
+
orthogonalizer = HEIWA_Orthogonalizer()
|
|
27
|
+
return heiwa, orthogonalizer
|
|
28
|
+
|
|
29
|
+
def load_thermo(hw_setting):
|
|
30
|
+
thermo = ThermoPred(hw_setting=hw_setting)
|
|
31
|
+
return thermo
|