CoolFace
Datasetpublic

si264/era-directed-evolution

Official repository for datasets and experimental results for "Efficient, Few-shot Directed Evolution with Energy Rank Alignment".

sourceHugging Facemitupdated 8mo agoView on Hugging Face
2likes603downloads
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