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
2likes832downloads
sample_esm_first_round.py193 linesDownload Raw Back to iterative_alignment_experiment_dpo
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