si264/era-directed-evolution
Official repository for datasets and experimental results for "Efficient, Few-shot Directed Evolution with Energy Rank Alignment".
2603
1import os2import torch3import numpy as np4import pandas as pd5import h5py6import re7from omegaconf import OmegaConf8import h5py9import lightning as L10from pera.nn import BidirectionalModel, sample_components_from_bidirectional_transformer, sample_perturbations, sample_embedding_perturbations11from esm.tokenization.sequence_tokenizer import EsmSequenceTokenizer12from Bio.Seq import Seq13 14device = torch.device("cuda:0")15 16sequence_tokenizer = EsmSequenceTokenizer()17 18import argparse19 20# set up parser21parser = parser = argparse.ArgumentParser(description="Calculating the log-likelihood of a sequence")22parser.add_argument('--target', type=str, required=True, help='Dataset as a string')23parser.add_argument('--num_samples', type=int, required=False, default=384, help='Number of samples to process (default: 100000)')24parser.add_argument('--alignment_round', type=int, required=False, default=1, help='Alignment round as an integer')25parser.add_argument('--version_number', type=str, required=False, default=1, help='Version number as a string')26parser.add_argument('--replicate', type=int, required=False, default=1, help='Replicate number as an integer')27args = parser.parse_args()28 29target = args.target30alignment_round = args.alignment_round31version_number = args.version_number32num_samples = args.num_samples33replicate = args.replicate34 35cfg_filename = f"{target}/lightning_logs_round_{alignment_round}/{version_number}/config.yaml"36network_filename = f"{target}/lightning_logs_round_{alignment_round}/{version_number}/checkpoints/best_model.ckpt"37save_folder_name = f"{target}/aligned_{alignment_round}_{num_samples}_{replicate}"38 39cfg = OmegaConf.load(cfg_filename)40sampling_temperature=141OmegaConf.update(cfg, "train.lightning_model_args.sampling_temperature", sampling_temperature)42esm_model = BidirectionalModel(cfg["nn"]["model"], 43 cfg["nn"]["model_args"],44 **cfg["train"]["lightning_model_args"]).to(device)45esm_model.load_model_from_ckpt(network_filename)46esm_model.eval()47print("")48mask_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["mask"]49bos_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["bos"]50eos_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["eos"]51pad_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["pad"]52 53 54os.makedirs(save_folder_name, exist_ok=True)55 56past_generations =[f"{target}/base_model_{num_samples}"]57for i in range(alignment_round):58 past_generations.append(f"{target}/aligned_{i}_{num_samples}_{replicate}")59 60previous_unmasked_sequences_decoded = []61 62for round in past_generations:63 trpb = torch.load(f"{round}/trpb_{replicate}.pt")64 previous_unmasked_sequences_decoded.extend(trpb['all_unmasked_sequences_decoded'])65 66data = target # "GB1", "ParD2", "TEV", "TrpB3F", "TrpB3I", "TrpB4"67data_root_path = "/scratch/groups/rotskoff/sebastian/era/protein_era/data"68 69sequence_tokenizer = EsmSequenceTokenizer()70 71if data.startswith("TrpB"):72 df = pd.read_csv(f"{data_root_path}/TrpB/scale2max/{data}.csv")73 with open(f"{data_root_path}/TrpB/TrpB.fasta", "r") as file:74 parent_sequence_decoded = file.readlines()[1].strip()75 76elif data == "DHFR":77 df = pd.read_csv(f"{data_root_path}/{data}/scale2max/{data}.csv")78 with open(f"{data_root_path}/{data}/{data}.fasta", "r") as file:79 nucleotide_seq = file.readlines()[1].strip()80 nucleotide_seq = Seq(nucleotide_seq)81 parent_sequence_decoded = str(nucleotide_seq.translate()) # Translate to amino acid sequence82 83else:84 df = pd.read_csv(f"{data_root_path}/{data}/scale2max/{data}.csv")85 with open(f"{data_root_path}/{data}/{data}.fasta", "r") as file:86 parent_sequence_decoded = file.readlines()[1].strip()87 88if data != "GB1": 89 muts = df["muts"].iloc[0]90else:91 muts = df["muts"].iloc[100000]92 93numbers = re.findall(r'\d+', muts)94mask_indices = list(map(int, numbers))95num_masks_per_sequence = num_samples // 496num_to_generate_per_mask = 497 98 99parent_sequence = torch.tensor(sequence_tokenizer.encode(parent_sequence_decoded, 100 add_special_tokens=True), device=device).unsqueeze(0).long()101sequence_length = parent_sequence.shape[1]102 103 104all_masked_sequences = []105all_unmasked_sequences_decoded = []106all_unmasked_sequences = []107all_logps = []108 109max_skips = 50110skips = 0111enforce_unique = True112 113while len(all_unmasked_sequences_decoded) < num_samples:114 115 print(len(all_unmasked_sequences_decoded))116 117 masked_sequences = parent_sequence.clone().repeat(num_to_generate_per_mask, 1)118 masked_sequences[:, mask_indices] = mask_token_sequence119 120 121 122 123 124 sequence_id = torch.ones((num_to_generate_per_mask, sequence_length), device=device).long() * 1125 126 structure_tokens = torch.ones((num_to_generate_per_mask, sequence_length), device=device).long() * 4096127 structure_tokens[:, 0] = 4098128 structure_tokens[:, -1] = 4097129 130 coords = torch.inf * torch.ones((num_to_generate_per_mask, sequence_length, 3, 3), device=device)131 132 average_plddt = torch.ones((num_to_generate_per_mask), device=device)133 134 per_res_plddt = torch.zeros((num_to_generate_per_mask, sequence_length), device=device)135 ss8_tokens = torch.zeros((num_to_generate_per_mask, sequence_length), device=device).long()136 sasa_tokens = torch.zeros((num_to_generate_per_mask, sequence_length), device=device).long()137 138 function_tokens = torch.zeros((num_to_generate_per_mask, sequence_length, 8), device=device).long()139 residue_annotation_tokens = torch.zeros((num_to_generate_per_mask, sequence_length, 16), device=device).long()140 141 142 143 144 with torch.no_grad():145 unmasked_sequences = sample_components_from_bidirectional_transformer(transformer_model=esm_model,146 masked_sequence_tokens=masked_sequences,147 structure_tokens=structure_tokens,148 average_plddt=average_plddt,149 per_res_plddt=per_res_plddt,150 ss8_tokens=ss8_tokens,151 sasa_tokens=sasa_tokens,152 function_tokens=function_tokens,153 residue_annotation_tokens=residue_annotation_tokens,154 bb_coords=coords,155 sequence_id=sequence_id,156 mask_token_sequence=mask_token_sequence,157 bos_token_sequence=bos_token_sequence,158 eos_token_sequence=eos_token_sequence,159 pad_token_sequence=pad_token_sequence,160 inference_batch_size=1)161 162 163 164 masked_indices = (masked_sequences == mask_token_sequence).float()165 logits = esm_model.nn(sequence_tokens=masked_sequences,166 structure_tokens=structure_tokens,167 average_plddt=average_plddt,168 per_res_plddt=per_res_plddt,169 ss8_tokens=ss8_tokens,170 sasa_tokens=sasa_tokens,171 function_tokens=function_tokens,172 residue_annotation_tokens=residue_annotation_tokens,173 sequence_id=sequence_id,174 bb_coords=coords)["sequence_logits"].detach()175 logps = torch.nn.functional.log_softmax(logits/sampling_temperature, dim=-1)176 logps = torch.gather(logps, dim=-1, index=unmasked_sequences.unsqueeze(-1)).squeeze(-1)177 logps = (logps * masked_indices).sum(-1).detach()178 179 decoded_seqs = [sequence.replace(" ", "") for sequence in sequence_tokenizer.batch_decode(unmasked_sequences[:, 1:-1])]180 for seq, logp, masked_seq, unmasked_seq in zip(decoded_seqs, logps, masked_sequences, unmasked_sequences):181 if enforce_unique and (seq in all_unmasked_sequences_decoded or seq in previous_unmasked_sequences_decoded):182 skips += 1183 if skips >= max_skips:184 enforce_unique=False185 continue186 else:187 skips=0188 all_unmasked_sequences_decoded.append(seq)189 all_logps.append(logp)190 all_masked_sequences.append(masked_seq)191 all_unmasked_sequences.append(unmasked_seq)192 193 194 195all_unmasked_sequences_decoded = all_unmasked_sequences_decoded[:num_samples]196all_masked_sequences = all_masked_sequences[:num_samples]197all_unmasked_sequences = all_unmasked_sequences[:num_samples]198all_logps = all_logps[:num_samples]199 200all_masked_sequences = torch.stack(all_masked_sequences, dim=0)201all_unmasked_sequences = torch.stack(all_unmasked_sequences, dim=0)202all_logps = torch.stack(all_logps, dim=0)203 204 205 206to_save = {"parent_sequence": parent_sequence,207 "all_masked_sequences": all_masked_sequences,208 "all_unmasked_sequences": all_unmasked_sequences,209 "all_unmasked_sequences_decoded": all_unmasked_sequences_decoded,210 "all_logps": all_logps}211torch.save(to_save, f"{save_folder_name}/trpb_{replicate}.pt")212 