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.
Files changed (55) hide show
  1. heiwa-0.1.0/PKG-INFO +13 -0
  2. heiwa-0.1.0/pyproject.toml +40 -0
  3. heiwa-0.1.0/setup.cfg +4 -0
  4. heiwa-0.1.0/src/heiwa/__init__.py +43 -0
  5. heiwa-0.1.0/src/heiwa/batch_run.py +67 -0
  6. heiwa-0.1.0/src/heiwa/generate/esm3_gen.py +159 -0
  7. heiwa-0.1.0/src/heiwa/generate/generate.py +105 -0
  8. heiwa-0.1.0/src/heiwa/generate/generate_utils.py +18 -0
  9. heiwa-0.1.0/src/heiwa/generate/modeller_script.py +89 -0
  10. heiwa-0.1.0/src/heiwa/hardware.py +24 -0
  11. heiwa-0.1.0/src/heiwa/loaders/dataloader.py +71 -0
  12. heiwa-0.1.0/src/heiwa/loaders/model_loader.py +31 -0
  13. heiwa-0.1.0/src/heiwa/loaders/tensor_loader.py +91 -0
  14. heiwa-0.1.0/src/heiwa/model_assets/ExposedResiduePrediction_epoch_final_model.pth.gz +0 -0
  15. heiwa-0.1.0/src/heiwa/model_assets/HEIWA_Hybridizer_epoch_final_model.pth.gz +0 -0
  16. heiwa-0.1.0/src/heiwa/model_assets/HEIWA_Imp_epoch_final_model.pth.gz +0 -0
  17. heiwa-0.1.0/src/heiwa/model_assets/TIM_epoch_final_model.pth.gz +0 -0
  18. heiwa-0.1.0/src/heiwa/model_assets/Thermo_epoch_final_model.pth.gz +0 -0
  19. heiwa-0.1.0/src/heiwa/model_assets/avg1536_wt.pt +0 -0
  20. heiwa-0.1.0/src/heiwa/model_assets/eigvec.pt +0 -0
  21. heiwa-0.1.0/src/heiwa/model_assets/logit_ranks.npy +0 -0
  22. heiwa-0.1.0/src/heiwa/model_assets/rmsf_ranks.npy +0 -0
  23. heiwa-0.1.0/src/heiwa/model_assets/sd1536_wt.pt +0 -0
  24. heiwa-0.1.0/src/heiwa/model_run/batch_accumulator.py +58 -0
  25. heiwa-0.1.0/src/heiwa/model_run/epoch_run.py +117 -0
  26. heiwa-0.1.0/src/heiwa/model_run/rmsf_convert.py +11 -0
  27. heiwa-0.1.0/src/heiwa/models/Water3P_Implicit.py +54 -0
  28. heiwa-0.1.0/src/heiwa/models/basic_layers.py +89 -0
  29. heiwa-0.1.0/src/heiwa/models/data_containers.py +8 -0
  30. heiwa-0.1.0/src/heiwa/models/file_handling.py +27 -0
  31. heiwa-0.1.0/src/heiwa/models/model.py +77 -0
  32. heiwa-0.1.0/src/heiwa/models/orthogonalizer.py +20 -0
  33. heiwa-0.1.0/src/heiwa/models/part_HEIWA_Hybridizer.py +41 -0
  34. heiwa-0.1.0/src/heiwa/models/part_HEIWA_Imp.py +139 -0
  35. heiwa-0.1.0/src/heiwa/models/part_SASA.py +34 -0
  36. heiwa-0.1.0/src/heiwa/models/part_TIM.py +52 -0
  37. heiwa-0.1.0/src/heiwa/models/part_thermo.py +44 -0
  38. heiwa-0.1.0/src/heiwa/predict.py +100 -0
  39. heiwa-0.1.0/src/heiwa/readers/read_cif.py +90 -0
  40. heiwa-0.1.0/src/heiwa/readers/read_dssp.py +76 -0
  41. heiwa-0.1.0/src/heiwa/readers/read_pdb.py +119 -0
  42. heiwa-0.1.0/src/heiwa/readers/read_utils.py +18 -0
  43. heiwa-0.1.0/src/heiwa/readers/reader.py +32 -0
  44. heiwa-0.1.0/src/heiwa/replace_esm3/blocks.py +164 -0
  45. heiwa-0.1.0/src/heiwa/replace_esm3/encoding.py +246 -0
  46. heiwa-0.1.0/src/heiwa/replace_esm3/esm3.py +652 -0
  47. heiwa-0.1.0/src/heiwa/replace_esm3/generation.py +845 -0
  48. heiwa-0.1.0/src/heiwa/replace_esm3/sequence_tokenizer.py +137 -0
  49. heiwa-0.1.0/src/heiwa/replace_esm3/transformer_stack.py +105 -0
  50. heiwa-0.1.0/src/heiwa/utils.py +42 -0
  51. heiwa-0.1.0/src/heiwa.egg-info/PKG-INFO +13 -0
  52. heiwa-0.1.0/src/heiwa.egg-info/SOURCES.txt +53 -0
  53. heiwa-0.1.0/src/heiwa.egg-info/dependency_links.txt +1 -0
  54. heiwa-0.1.0/src/heiwa.egg-info/requires.txt +10 -0
  55. 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,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+
@@ -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