si264/era-directed-evolution
Official repository for datasets and experimental results for "Efficient, Few-shot Directed Evolution with Energy Rank Alignment".
2667
1import torch2import re3import pandas as pd4import numpy as np5import matplotlib.pyplot as plt6import h5py7from omegaconf import OmegaConf8from esm.tokenization.sequence_tokenizer import EsmSequenceTokenizer9from Bio.PDB import PDBList, PDBParser, is_aa10 11device = torch.device("cuda:0")12 13# Optional: map 3-letter residue names to 1-letter codes14three_to_one = {15 'ALA': 'A', 'ARG': 'R', 'ASN': 'N', 'ASP': 'D',16 'CYS': 'C', 'GLN': 'Q', 'GLU': 'E', 'GLY': 'G',17 'HIS': 'H', 'ILE': 'I', 'LEU': 'L', 'LYS': 'K',18 'MET': 'M', 'PHE': 'F', 'PRO': 'P', 'SER': 'S',19 'THR': 'T', 'TRP': 'W', 'TYR': 'Y', 'VAL': 'V',20 'SEC': 'U', 'PYL': 'O', 'ASX': 'B', 'GLX': 'Z',21 'XLE': 'J', 'UNK': 'X'22}23 24def get_backbone_coords_from_local_pdb(pdb_path, chain_id='A', sequence_length=None, target="data", device=device):25 """26 Load backbone coordinates and residue types from a local PDB file.27 28 Returns:29 coords_tensor: torch.Tensor of shape (1, N, 3, 3)30 residue_types: List of one-letter residue codes31 """32 parser = PDBParser(QUIET=True)33 structure = parser.get_structure("local_structure", pdb_path)34 35 coords = []36 residue_types = []37 model = structure[0]38 39 if chain_id not in model:40 raise ValueError(f"Chain {chain_id} not found in {pdb_path}")41 42 chain = model[chain_id]43 44 for residue in chain:45 if sequence_length is not None and len(coords) >= sequence_length:46 break47 if not is_aa(residue):48 continue49 try:50 n = residue['N'].get_coord()51 ca = residue['CA'].get_coord()52 c = residue['C'].get_coord()53 coords.append([n, ca, c])54 resname = residue.get_resname().upper()55 residue_types.append(three_to_one.get(resname, 'X')) # default to 'X' if unknown56 except KeyError:57 continue58 59 if not coords:60 raise ValueError("No residues with complete backbone atoms found.")61 62 # Add infinity-padding before and after63 pad = [[float('inf')]*3, [float('inf')]*3, [float('inf')]*3]64 coords.insert(0, pad)65 coords.append(pad)66 67 if target == "ParD2":68 coords = [pad, pad] + coords + [pad, pad]69 elif target == "ParD3":70 coords = [pad]*2 + coords + [pad]*671 elif target == "TrpB4":72 coords = [pad] + coords73 74 coords_tensor = torch.tensor(coords, device=device).unsqueeze(0) # (1, N, 3, 3)75 76 return coords_tensor, residue_types77 78num_replicates = 1079campaign_number = 2 # change this according to the campaign we are interested in80dataset_size = 96 # change this according to the dataset size we are interested in81 82sequence_tokenizer = EsmSequenceTokenizer()83 84datasets = ["GB1", "TrpB4"]85data_root_path = "/global/cfs/projectdirs/m4235/sebastian/data"86 87for data in datasets:88 print(data)89 for i in range(num_replicates):90 cfg_filename = f"./config.yaml"91 cfg = OmegaConf.load(cfg_filename)92 sampling_temperature=193 OmegaConf.update(cfg, "train.lightning_model_args.sampling_temperature", sampling_temperature)94 mask_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["mask"]95 bos_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["bos"]96 eos_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["eos"]97 pad_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["pad"]98 99 if not data.startswith("TrpB"):100 df = pd.read_csv(f"{data_root_path}/{data}/scale2max/{data}.csv")101 with open(f"{data_root_path}/{data}/{data}.fasta", "r") as file:102 parent_sequence_decoded = file.readlines()[1].strip()103 104 else:105 df = pd.read_csv(f"{data_root_path}/TrpB/scale2max/{data}.csv")106 with open(f"{data_root_path}/TrpB/TrpB.fasta", "r") as file:107 parent_sequence_decoded = file.readlines()[1].strip()108 109 if data != "GB1": 110 muts = df["muts"].iloc[0]111 else:112 muts = df["muts"].iloc[100000]113 114 numbers = re.findall(r'\d+', muts)115 mask_indices = list(map(int, numbers))116 # mask_indices = [i-1 for i in mask_indices] #convert to 0-based indexing117 118 fitness_scores = []119 120# Load from base_model_{dataset_size}121 trpb_base = torch.load(f"./{data}/base_model_{dataset_size}/trpb_post_rd_{campaign_number-1}_{i}.pt")122 all_unmasked_sequences_decoded_base = trpb_base["all_unmasked_sequences_decoded"]123 all_unmasked_sequences_base = trpb_base["all_unmasked_sequences"]124 all_masked_sequences_base = trpb_base["all_masked_sequences"]125 all_unmasked_sequences_base = all_unmasked_sequences_base.reshape(-1, all_unmasked_sequences_base.shape[-1])126 all_logps_base = trpb_base["all_logps"]127 128 for unmasked_sequence_decoded, unmasked_sequence in zip(all_unmasked_sequences_decoded_base, all_unmasked_sequences_base):129 index_residue_0 = unmasked_sequence_decoded[mask_indices[0]-1]130 index_residue_1 = unmasked_sequence_decoded[mask_indices[1]-1]131 index_residue_2 = unmasked_sequence_decoded[mask_indices[2]-1]132 try:133 index_residue_3 = unmasked_sequence_decoded[mask_indices[3]-1]134 mutations = [index_residue_0, index_residue_1, index_residue_2, index_residue_3]135 muts = ''.join(mutations)136 except:137 mutations = [index_residue_0, index_residue_1, index_residue_2]138 muts = ''.join(mutations)139 140 df_filtered = df[df["AAs"] == muts]141 142 if len(df_filtered) == 0:143 if torch.any((unmasked_sequence[1:-1] > 23) | (unmasked_sequence[1:-1] < 4)):144 print(f"Invalid sequence {muts}")145 fitness_score = -2146 else:147 print(f"Invalid sequence {muts}")148 fitness_score = -2149 else:150 fitness_score = df_filtered["fitness"].values[0]151 fitness_scores.append(fitness_score)152 153 # Load from aligned_0_{dataset_size}154 trpb_aligned = torch.load(f"./{data}/aligned_{campaign_number-2}_{dataset_size}_{i}/trpb_post_rd_{campaign_number-1}_{i}.pt")155 all_unmasked_sequences_decoded_aligned_0 = trpb_aligned["all_unmasked_sequences_decoded"]156 all_unmasked_sequences_aligned_0 = trpb_aligned["all_unmasked_sequences"]157 all_masked_sequences_aligned_0 = trpb_aligned["all_masked_sequences"]158 all_unmasked_sequences_aligned_0 = all_unmasked_sequences_aligned_0.reshape(-1, all_unmasked_sequences_aligned_0.shape[-1])159 all_logps_aligned_0 = trpb_aligned["all_logps"]160 161 for unmasked_sequence_decoded, unmasked_sequence in zip(all_unmasked_sequences_decoded_aligned_0, all_unmasked_sequences_aligned_0):162 index_residue_0 = unmasked_sequence_decoded[mask_indices[0]-1]163 index_residue_1 = unmasked_sequence_decoded[mask_indices[1]-1]164 index_residue_2 = unmasked_sequence_decoded[mask_indices[2]-1]165 try:166 index_residue_3 = unmasked_sequence_decoded[mask_indices[3]-1]167 mutations = [index_residue_0, index_residue_1, index_residue_2, index_residue_3]168 muts = ''.join(mutations)169 except:170 mutations = [index_residue_0, index_residue_1, index_residue_2]171 muts = ''.join(mutations)172 173 df_filtered = df[df["AAs"] == muts]174 175 if len(df_filtered) == 0:176 if torch.any((unmasked_sequence[1:-1] > 23) | (unmasked_sequence[1:-1] < 4)):177 print(f"Invalid sequence {muts}")178 fitness_score = -2179 else:180 print(f"Invalid sequence {muts}")181 fitness_score = -2182 else:183 fitness_score = df_filtered["fitness"].values[0]184 fitness_scores.append(fitness_score)185 186 # Load from aligned_1_{dataset_size}187 trpb_aligned_1 = torch.load(f"./{data}/aligned_{campaign_number-1}_{dataset_size}_{i}/trpb_{i}.pt")188 all_unmasked_sequences_decoded_aligned_1 = trpb_aligned_1["all_unmasked_sequences_decoded"]189 all_unmasked_sequences_aligned_1 = trpb_aligned_1["all_unmasked_sequences"]190 all_masked_sequences_aligned_1 = trpb_aligned_1["all_masked_sequences"]191 all_unmasked_sequences_aligned_1 = all_unmasked_sequences_aligned_1.reshape(-1, all_unmasked_sequences_aligned_1.shape[-1])192 all_logps_aligned_1 = trpb_aligned_1["all_logps"]193 194 for unmasked_sequence_decoded, unmasked_sequence in zip(all_unmasked_sequences_decoded_aligned_1, all_unmasked_sequences_aligned_1):195 index_residue_0 = unmasked_sequence_decoded[mask_indices[0]-1]196 index_residue_1 = unmasked_sequence_decoded[mask_indices[1]-1]197 index_residue_2 = unmasked_sequence_decoded[mask_indices[2]-1]198 try:199 index_residue_3 = unmasked_sequence_decoded[mask_indices[3]-1]200 mutations = [index_residue_0, index_residue_1, index_residue_2, index_residue_3]201 muts = ''.join(mutations)202 except:203 mutations = [index_residue_0, index_residue_1, index_residue_2]204 muts = ''.join(mutations)205 206 df_filtered = df[df["AAs"] == muts]207 208 if len(df_filtered) == 0:209 if torch.any((unmasked_sequence[1:-1] > 23) | (unmasked_sequence[1:-1] < 4)):210 print(f"Invalid sequence {muts}")211 fitness_score = -2212 else:213 print(f"Invalid sequence {muts}")214 fitness_score = -2215 else:216 fitness_score = df_filtered["fitness"].values[0]217 fitness_scores.append(fitness_score)218 219 # Concatenate the sequences and logps from all models220 all_unmasked_sequences = torch.cat((all_unmasked_sequences_base, all_unmasked_sequences_aligned_0, all_unmasked_sequences_aligned_1),dim=0)221 all_masked_sequences = torch.cat((all_masked_sequences_base, all_masked_sequences_aligned_0, all_masked_sequences_aligned_1),dim=0)222 print(all_logps_base.shape, all_logps_aligned_0.shape, all_logps_aligned_1.shape)223 all_logps = torch.cat((all_logps_base, all_logps_aligned_0, all_logps_aligned_1),dim=0)224 225 all_fitness_scores = fitness_scores226 227 # Check for duplicates in all_unmasked_sequences228 unique_sequences, counts = torch.unique(all_unmasked_sequences, dim=0, return_counts=True)229 num_duplicates = torch.sum(counts > 1).item()230 print(f"Number of duplicate sequences: {num_duplicates}")231 232 all_fitness_scores = np.array(all_fitness_scores)233 all_fitness_scores = np.where(all_fitness_scores > 0, -np.log(all_fitness_scores), 10)234 235 sampling_temperature = 1 # hard-coding a sampling temperature of 1 for mixed-temperature alignment236 237 sequence_length = all_unmasked_sequences.shape[1]238 239 sequence_id = torch.ones((all_unmasked_sequences.shape[0], sequence_length), device=device).long() * 1240 241 structure_tokens = torch.ones((1, sequence_length), device=device).long() * 4096242 structure_tokens[:, 0] = 4098243 structure_tokens[:, -1] = 4097244 245 coords, residue_types = get_backbone_coords_from_local_pdb(f"{data_root_path}/{data}/{data}.pdb", chain_id='A', sequence_length=sequence_length-2, target=data) if not data.startswith("TrpB") else get_backbone_coords_from_local_pdb(f"{data_root_path}/TrpB/TrpB.pdb", chain_id='A', sequence_length=sequence_length-2, target=data)246 247 # parent sequence sanity check248 coords_trimmed = coords[:, 1:-1] # shape: (1, N-2, 3, 3)249 250 # Step 2: Determine mask of non-padding residues (i.e., not all coords are inf)251 valid_mask = ~(torch.isinf(coords_trimmed).view(-1, 9).any(dim=1)) # shape: (N-2,)252 residues_to_compare = [r for r, valid in zip(list(parent_sequence_decoded), valid_mask) if valid]253 254 if residue_types != residues_to_compare:255 print("Residue mismatch detected!")256 for i, (ref, pdb) in enumerate(zip(residues_to_compare, residue_types)):257 if ref != pdb:258 print(f"Position {i}: expected {ref}, got {pdb}")259 else:260 print("Residues match.")261 print(coords.shape)262 263 assert coords.shape[1] == sequence_length, f"Coords length {coords.shape[1]} does not match sequence length {sequence_length}"264 265 average_plddt = torch.ones((1), device=device)266 267 per_res_plddt = torch.zeros((1, sequence_length), device=device)268 ss8_tokens = torch.zeros((1, sequence_length), device=device).long()269 sasa_tokens = torch.zeros((1, sequence_length), device=device).long()270 271 function_tokens = torch.zeros((1, sequence_length, 8), device=device).long()272 residue_annotation_tokens = torch.zeros((1, sequence_length, 16), device=device).long()273 274 275 with h5py.File(f"./{data}/alignment_dataset_{campaign_number}_{dataset_size}_from_ESM3_{i}.hdf5", "w") as f:276 masked_sequence_tokens = f.create_dataset("masked_sequence_tokens", data=all_masked_sequences.cpu().numpy())277 unmasked_sequence_tokens = f.create_dataset("unmasked_sequence_tokens", data=all_unmasked_sequences.cpu().numpy())278 sequence_id = f.create_dataset("sequence_id", data=sequence_id.cpu().numpy())279 structure_tokens = f.create_dataset("structural_tokens", data=structure_tokens.cpu().numpy())280 coords = f.create_dataset("bb_coords", data=coords.cpu().numpy())281 average_plddt = f.create_dataset("average_plddt", data=average_plddt.cpu().numpy())282 per_res_plddt = f.create_dataset("per_res_plddt", data=per_res_plddt.cpu().numpy())283 ss8_tokens = f.create_dataset("ss8_tokens", data=ss8_tokens.cpu().numpy())284 sasa_tokens = f.create_dataset("sasa_tokens", data=sasa_tokens.cpu().numpy())285 function_tokens = f.create_dataset("function_tokens", data=function_tokens.cpu().numpy())286 residue_annotation_tokens = f.create_dataset("residue_annotation_tokens", data=residue_annotation_tokens.cpu().numpy())287 288 ref_logps = f.create_dataset("ref_logps", data=all_logps.cpu().numpy())289 energies = f.create_dataset("energies", data=all_fitness_scores)290 291 292 f.attrs["num_prompts"] = 1293 f.attrs["num_examples_per_prompt"] = masked_sequence_tokens.shape[0]294 f.attrs["fixed_bb_coords"] = True295 f.attrs["fixed_average_plddt"] = True296 f.attrs["fixed_per_res_plddt"] = True297 f.attrs["fixed_ss8_tokens"] = True298 f.attrs["fixed_sasa_tokens"] = True299 f.attrs["fixed_function_tokens"] = True300 f.attrs["fixed_residue_annotation_tokens"] = True301 f.attrs["fixed_structural_tokens"] = True302 f.attrs["sampling_temperature"] = sampling_temperature303 304 