si264/era-directed-evolution
Official repository for datasets and experimental results for "Efficient, Few-shot Directed Evolution with Energy Rank Alignment".
2832
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('--version_number', type=str, required=False, default=1, help='Version number as a string')25parser.add_argument('--replicate', type=int, required=False, default=1, help='Replicate number as an integer')26args = parser.parse_args()27 28target = args.target29num_samples = args.num_samples30replicate = args.replicate31 32cfg_filename = "./config.yaml"33network_filename = "/scratch/groups/rotskoff/sebastian/era/protein_era/models/esm3clm/esm3_clm.pt"34save_folder_name = f"{target}/base_model_{num_samples}"35 36cfg = OmegaConf.load(cfg_filename)37sampling_temperature=138OmegaConf.update(cfg, "train.lightning_model_args.sampling_temperature", sampling_temperature)39OmegaConf.update(cfg, "train.lightning_model_args.better_energy", "lower")40esm_model = BidirectionalModel(cfg["nn"]["model"], 41 cfg["nn"]["model_args"],42 **cfg["train"]["lightning_model_args"]).to(device)43esm_model.load_model_from_ckpt(network_filename)44esm_model.eval()45print("")46mask_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["mask"]47bos_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["bos"]48eos_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["eos"]49pad_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["pad"]50 51 52os.makedirs(save_folder_name, exist_ok=True)53 54data = target # "GB1", "ParD2", "TEV", "TrpB3F", "TrpB3I", "TrpB4"55data_root_path = "/scratch/groups/rotskoff/sebastian/era/protein_era/data"56 57sequence_tokenizer = EsmSequenceTokenizer()58 59if data.startswith("TrpB"):60 df = pd.read_csv(f"{data_root_path}/TrpB/scale2max/{data}.csv")61 with open(f"{data_root_path}/TrpB/TrpB.fasta", "r") as file:62 parent_sequence_decoded = file.readlines()[1].strip()63 64elif data == "DHFR":65 df = pd.read_csv(f"{data_root_path}/{data}/scale2max/{data}.csv")66 with open(f"{data_root_path}/{data}/{data}.fasta", "r") as file:67 nucleotide_seq = file.readlines()[1].strip()68 nucleotide_seq = Seq(nucleotide_seq)69 parent_sequence_decoded = str(nucleotide_seq.translate()) # Translate to amino acid sequence70 71else:72 df = pd.read_csv(f"{data_root_path}/{data}/scale2max/{data}.csv")73 with open(f"{data_root_path}/{data}/{data}.fasta", "r") as file:74 parent_sequence_decoded = file.readlines()[1].strip()75 76if data != "GB1": 77 muts = df["muts"].iloc[0]78else:79 muts = df["muts"].iloc[100000]80 81numbers = re.findall(r'\d+', muts)82mask_indices = list(map(int, numbers))83num_masks_per_sequence = num_samples // 484num_to_generate_per_mask = 485 86 87parent_sequence = torch.tensor(sequence_tokenizer.encode(parent_sequence_decoded, 88 add_special_tokens=True), device=device).unsqueeze(0).long()89sequence_length = parent_sequence.shape[1]90 91 92all_masked_sequences = []93all_unmasked_sequences_decoded = []94all_unmasked_sequences = []95all_logps = []96 97 98while len(all_unmasked_sequences_decoded) < num_samples:99 100 print(len(all_unmasked_sequences_decoded))101 102 masked_sequences = parent_sequence.clone().repeat(num_to_generate_per_mask, 1)103 masked_sequences[:, mask_indices] = mask_token_sequence104 105 106 107 108 109 sequence_id = torch.ones((num_to_generate_per_mask, sequence_length), device=device).long() * 1110 111 structure_tokens = torch.ones((num_to_generate_per_mask, sequence_length), device=device).long() * 4096112 structure_tokens[:, 0] = 4098113 structure_tokens[:, -1] = 4097114 115 coords = torch.inf * torch.ones((num_to_generate_per_mask, sequence_length, 3, 3), device=device)116 117 average_plddt = torch.ones((num_to_generate_per_mask), device=device)118 119 per_res_plddt = torch.zeros((num_to_generate_per_mask, sequence_length), device=device)120 ss8_tokens = torch.zeros((num_to_generate_per_mask, sequence_length), device=device).long()121 sasa_tokens = torch.zeros((num_to_generate_per_mask, sequence_length), device=device).long()122 123 function_tokens = torch.zeros((num_to_generate_per_mask, sequence_length, 8), device=device).long()124 residue_annotation_tokens = torch.zeros((num_to_generate_per_mask, sequence_length, 16), device=device).long()125 126 127 128 129 with torch.no_grad():130 unmasked_sequences = sample_components_from_bidirectional_transformer(transformer_model=esm_model,131 masked_sequence_tokens=masked_sequences,132 structure_tokens=structure_tokens,133 average_plddt=average_plddt,134 per_res_plddt=per_res_plddt,135 ss8_tokens=ss8_tokens,136 sasa_tokens=sasa_tokens,137 function_tokens=function_tokens,138 residue_annotation_tokens=residue_annotation_tokens,139 bb_coords=coords,140 sequence_id=sequence_id,141 mask_token_sequence=mask_token_sequence,142 bos_token_sequence=bos_token_sequence,143 eos_token_sequence=eos_token_sequence,144 pad_token_sequence=pad_token_sequence,145 inference_batch_size=1)146 147 148 149 masked_indices = (masked_sequences == mask_token_sequence).float()150 logits = esm_model.nn(sequence_tokens=masked_sequences,151 structure_tokens=structure_tokens,152 average_plddt=average_plddt,153 per_res_plddt=per_res_plddt,154 ss8_tokens=ss8_tokens,155 sasa_tokens=sasa_tokens,156 function_tokens=function_tokens,157 residue_annotation_tokens=residue_annotation_tokens,158 sequence_id=sequence_id,159 bb_coords=coords)["sequence_logits"].detach()160 logps = torch.nn.functional.log_softmax(logits/sampling_temperature, dim=-1)161 logps = torch.gather(logps, dim=-1, index=unmasked_sequences.unsqueeze(-1)).squeeze(-1)162 logps = (logps * masked_indices).sum(-1).detach()163 164 decoded_seqs = [sequence.replace(" ", "") for sequence in sequence_tokenizer.batch_decode(unmasked_sequences[:, 1:-1])]165 for seq, logp, masked_seq, unmasked_seq in zip(decoded_seqs, logps, masked_sequences, unmasked_sequences):166 if seq in all_unmasked_sequences_decoded:167 continue168 else:169 all_unmasked_sequences_decoded.append(seq)170 all_logps.append(logp)171 all_masked_sequences.append(masked_seq)172 all_unmasked_sequences.append(unmasked_seq)173 174 175 176all_unmasked_sequences_decoded = all_unmasked_sequences_decoded[:num_samples]177all_masked_sequences = all_masked_sequences[:num_samples]178all_unmasked_sequences = all_unmasked_sequences[:num_samples]179all_logps = all_logps[:num_samples]180 181all_masked_sequences = torch.stack(all_masked_sequences, dim=0)182all_unmasked_sequences = torch.stack(all_unmasked_sequences, dim=0)183all_logps = torch.stack(all_logps, dim=0)184 185 186 187to_save = {"parent_sequence": parent_sequence,188 "all_masked_sequences": all_masked_sequences,189 "all_unmasked_sequences": all_unmasked_sequences,190 "all_unmasked_sequences_decoded": all_unmasked_sequences_decoded,191 "all_logps": all_logps}192torch.save(to_save, f"{save_folder_name}/trpb_{replicate}.pt")193 