si264/era-directed-evolution
Official repository for datasets and experimental results for "Efficient, Few-shot Directed Evolution with Energy Rank Alignment".
2575
1import argparse2import torch3import re4import pandas as pd5import numpy as np6import matplotlib.pyplot as plt7import h5py8from omegaconf import OmegaConf9from esm.tokenization.sequence_tokenizer import EsmSequenceTokenizer10from Bio.Seq import Seq11 12device = torch.device("cuda:0")13 14num_replicates = 1015campaign_number = 1 # change this according to the campaign we are interested in16dataset_size = 96 # change this according to the dataset size we are interested in17 18sequence_tokenizer = EsmSequenceTokenizer()19 20parser = argparse.ArgumentParser(description="Calculating the log-likelihood of a sequence")21parser.add_argument('--target', type=str, required=True, help='Dataset as a string')22args = parser.parse_args()23data = args.target24 25data_root_path = "/scratch/groups/rotskoff/sebastian/era/protein_era/data"26 27 28print(data)29for i in range(num_replicates):30 cfg_filename = f"./config.yaml"31 cfg = OmegaConf.load(cfg_filename)32 sampling_temperature=133 OmegaConf.update(cfg, "train.lightning_model_args.sampling_temperature", sampling_temperature)34 mask_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["mask"]35 bos_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["bos"]36 eos_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["eos"]37 pad_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["pad"]38 39 if data.startswith("TrpB"):40 df = pd.read_csv(f"{data_root_path}/TrpB/scale2max/{data}.csv")41 with open(f"{data_root_path}/TrpB/TrpB.fasta", "r") as file:42 parent_sequence_decoded = file.readlines()[1].strip()43 44 elif data == "DHFR":45 df = pd.read_csv(f"{data_root_path}/{data}/scale2max/{data}.csv")46 with open(f"{data_root_path}/{data}/{data}.fasta", "r") as file:47 nucleotide_seq = file.readlines()[1].strip()48 nucleotide_seq = Seq(nucleotide_seq)49 parent_sequence_decoded = str(nucleotide_seq.translate()) # Translate to amino acid sequence50 51 else:52 df = pd.read_csv(f"{data_root_path}/{data}/scale2max/{data}.csv")53 with open(f"{data_root_path}/{data}/{data}.fasta", "r") as file:54 parent_sequence_decoded = file.readlines()[1].strip()55 56 if data != "GB1": 57 muts = df["muts"].iloc[0]58 else:59 muts = df["muts"].iloc[100000]60 61 numbers = re.findall(r'\d+', muts)62 mask_indices = list(map(int, numbers))63 # mask_indices = [i-1 for i in mask_indices] #convert to 0-based indexing64 65 fitness_scores = []66 67 # Load from base_model_{dataset_size}68 trpb_base = torch.load(f"./{data}/base_model_{dataset_size}/trpb_post_rd_{campaign_number-1}_{i}.pt")69 all_unmasked_sequences_decoded_base = trpb_base["all_unmasked_sequences_decoded"]70 all_unmasked_sequences_base = trpb_base["all_unmasked_sequences"]71 all_masked_sequences_base = trpb_base["all_masked_sequences"]72 all_unmasked_sequences_base = all_unmasked_sequences_base.reshape(-1, all_unmasked_sequences_base.shape[-1])73 all_logps_base = trpb_base["all_logps"]74 75 for unmasked_sequence_decoded, unmasked_sequence in zip(all_unmasked_sequences_decoded_base, all_unmasked_sequences_base):76 index_residue_0 = unmasked_sequence_decoded[mask_indices[0]-1]77 index_residue_1 = unmasked_sequence_decoded[mask_indices[1]-1]78 index_residue_2 = unmasked_sequence_decoded[mask_indices[2]-1]79 try:80 index_residue_3 = unmasked_sequence_decoded[mask_indices[3]-1]81 mutations = [index_residue_0, index_residue_1, index_residue_2, index_residue_3]82 muts = ''.join(mutations)83 except:84 mutations = [index_residue_0, index_residue_1, index_residue_2]85 muts = ''.join(mutations)86 87 df_filtered = df[df["AAs"] == muts]88 89 if len(df_filtered) == 0:90 if torch.any((unmasked_sequence[1:-1] > 23) | (unmasked_sequence[1:-1] < 4)):91 print(f"Invalid sequence {muts}")92 fitness_score = -293 else:94 print(f"Invalid sequence {muts}")95 fitness_score = -296 else:97 fitness_score = df_filtered["fitness"].values[0]98 fitness_scores.append(fitness_score)99 100 # Load from aligned_0_{dataset_size}101 trpb_aligned = torch.load(f"./{data}/aligned_{campaign_number-1}_{dataset_size}_{i}/trpb_{i}.pt")102 all_unmasked_sequences_decoded_aligned_0 = trpb_aligned["all_unmasked_sequences_decoded"]103 all_unmasked_sequences_aligned_0 = trpb_aligned["all_unmasked_sequences"]104 all_masked_sequences_aligned_0 = trpb_aligned["all_masked_sequences"]105 all_unmasked_sequences_aligned_0 = all_unmasked_sequences_aligned_0.reshape(-1, all_unmasked_sequences_aligned_0.shape[-1])106 all_logps_aligned_0 = trpb_aligned["all_logps"]107 108 for unmasked_sequence_decoded, unmasked_sequence in zip(all_unmasked_sequences_decoded_aligned_0, all_unmasked_sequences_aligned_0):109 index_residue_0 = unmasked_sequence_decoded[mask_indices[0]-1]110 index_residue_1 = unmasked_sequence_decoded[mask_indices[1]-1]111 index_residue_2 = unmasked_sequence_decoded[mask_indices[2]-1]112 try:113 index_residue_3 = unmasked_sequence_decoded[mask_indices[3]-1]114 mutations = [index_residue_0, index_residue_1, index_residue_2, index_residue_3]115 muts = ''.join(mutations)116 except:117 mutations = [index_residue_0, index_residue_1, index_residue_2]118 muts = ''.join(mutations)119 120 df_filtered = df[df["AAs"] == muts]121 122 if len(df_filtered) == 0:123 if torch.any((unmasked_sequence[1:-1] > 23) | (unmasked_sequence[1:-1] < 4)):124 print(f"Invalid sequence {muts}")125 fitness_score = -2126 else:127 print(f"Invalid sequence {muts}")128 fitness_score = -2129 else:130 fitness_score = df_filtered["fitness"].values[0]131 fitness_scores.append(fitness_score)132 133 134 # Concatenate the sequences and logps from all models135 all_unmasked_sequences = torch.cat((all_unmasked_sequences_base, all_unmasked_sequences_aligned_0),dim=0)#136 all_masked_sequences = torch.cat((all_masked_sequences_base, all_masked_sequences_aligned_0),dim=0)#137 print(all_logps_base.shape, all_logps_aligned_0.shape)#138 all_logps = torch.cat((all_logps_base, all_logps_aligned_0),dim=0)#139 140 all_fitness_scores = fitness_scores141 142 # Check for duplicates in all_unmasked_sequences143 unique_sequences, counts = torch.unique(all_unmasked_sequences, dim=0, return_counts=True)144 num_duplicates = torch.sum(counts > 1).item()145 print(f"Number of duplicate sequences: {num_duplicates}")146 147 all_fitness_scores = np.array(all_fitness_scores)148 all_fitness_scores = np.where(all_fitness_scores > 0, -np.log(all_fitness_scores), 10)149 150 sampling_temperature = 1 # hard-coding a sampling temperature of 1 for mixed-temperature alignment151 152 sequence_length = all_unmasked_sequences.shape[1]153 154 sequence_id = torch.ones((all_unmasked_sequences.shape[0], sequence_length), device=device).long() * 1155 156 structure_tokens = torch.ones((1, sequence_length), device=device).long() * 4096157 structure_tokens[:, 0] = 4098158 structure_tokens[:, -1] = 4097159 160 coords = torch.inf * torch.ones((1, sequence_length, 3, 3), device=device)161 162 average_plddt = torch.ones((1), device=device)163 164 per_res_plddt = torch.zeros((1, sequence_length), device=device)165 ss8_tokens = torch.zeros((1, sequence_length), device=device).long()166 sasa_tokens = torch.zeros((1, sequence_length), device=device).long()167 168 function_tokens = torch.zeros((1, sequence_length, 8), device=device).long()169 residue_annotation_tokens = torch.zeros((1, sequence_length, 16), device=device).long()170 171 172 with h5py.File(f"./{data}/alignment_dataset_{campaign_number}_{dataset_size}_from_ESM3_{i}.hdf5", "w") as f:173 masked_sequence_tokens = f.create_dataset("masked_sequence_tokens", data=all_masked_sequences.cpu().numpy())174 unmasked_sequence_tokens = f.create_dataset("unmasked_sequence_tokens", data=all_unmasked_sequences.cpu().numpy())175 sequence_id = f.create_dataset("sequence_id", data=sequence_id.cpu().numpy())176 structure_tokens = f.create_dataset("structural_tokens", data=structure_tokens.cpu().numpy())177 coords = f.create_dataset("bb_coords", data=coords.cpu().numpy())178 average_plddt = f.create_dataset("average_plddt", data=average_plddt.cpu().numpy())179 per_res_plddt = f.create_dataset("per_res_plddt", data=per_res_plddt.cpu().numpy())180 ss8_tokens = f.create_dataset("ss8_tokens", data=ss8_tokens.cpu().numpy())181 sasa_tokens = f.create_dataset("sasa_tokens", data=sasa_tokens.cpu().numpy())182 function_tokens = f.create_dataset("function_tokens", data=function_tokens.cpu().numpy())183 residue_annotation_tokens = f.create_dataset("residue_annotation_tokens", data=residue_annotation_tokens.cpu().numpy())184 185 ref_logps = f.create_dataset("ref_logps", data=all_logps.cpu().numpy())186 energies = f.create_dataset("energies", data=all_fitness_scores)187 188 189 f.attrs["num_prompts"] = 1190 f.attrs["num_examples_per_prompt"] = masked_sequence_tokens.shape[0]191 f.attrs["fixed_bb_coords"] = True192 f.attrs["fixed_average_plddt"] = True193 f.attrs["fixed_per_res_plddt"] = True194 f.attrs["fixed_ss8_tokens"] = True195 f.attrs["fixed_sasa_tokens"] = True196 f.attrs["fixed_function_tokens"] = True197 f.attrs["fixed_residue_annotation_tokens"] = True198 f.attrs["fixed_structural_tokens"] = True199 f.attrs["sampling_temperature"] = sampling_temperature200 