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
2likes575downloads
make_alignment_dataset_second_round.py200 linesDownload Raw Back to iterative_alignment_experiment_dpo
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