Team Ai
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
2likes667downloads
create_alignment_dataset_third_round.py304 linesDownload Raw Back to iterative_alignment_experiment_structure
1import torch2import re3import pandas as pd4import numpy as np5import matplotlib.pyplot as plt6import h5py7from omegaconf import OmegaConf8from esm.tokenization.sequence_tokenizer import EsmSequenceTokenizer9from Bio.PDB import PDBList, PDBParser, is_aa10 11device = torch.device("cuda:0")12 13# Optional: map 3-letter residue names to 1-letter codes14three_to_one = {15    'ALA': 'A', 'ARG': 'R', 'ASN': 'N', 'ASP': 'D',16    'CYS': 'C', 'GLN': 'Q', 'GLU': 'E', 'GLY': 'G',17    'HIS': 'H', 'ILE': 'I', 'LEU': 'L', 'LYS': 'K',18    'MET': 'M', 'PHE': 'F', 'PRO': 'P', 'SER': 'S',19    'THR': 'T', 'TRP': 'W', 'TYR': 'Y', 'VAL': 'V',20    'SEC': 'U', 'PYL': 'O', 'ASX': 'B', 'GLX': 'Z',21    'XLE': 'J', 'UNK': 'X'22}23 24def get_backbone_coords_from_local_pdb(pdb_path, chain_id='A', sequence_length=None, target="data", device=device):25    """26    Load backbone coordinates and residue types from a local PDB file.27 28    Returns:29        coords_tensor: torch.Tensor of shape (1, N, 3, 3)30        residue_types: List of one-letter residue codes31    """32    parser = PDBParser(QUIET=True)33    structure = parser.get_structure("local_structure", pdb_path)34 35    coords = []36    residue_types = []37    model = structure[0]38 39    if chain_id not in model:40        raise ValueError(f"Chain {chain_id} not found in {pdb_path}")41 42    chain = model[chain_id]43 44    for residue in chain:45        if sequence_length is not None and len(coords) >= sequence_length:46            break47        if not is_aa(residue):48            continue49        try:50            n = residue['N'].get_coord()51            ca = residue['CA'].get_coord()52            c = residue['C'].get_coord()53            coords.append([n, ca, c])54            resname = residue.get_resname().upper()55            residue_types.append(three_to_one.get(resname, 'X'))  # default to 'X' if unknown56        except KeyError:57            continue58 59    if not coords:60        raise ValueError("No residues with complete backbone atoms found.")61 62    # Add infinity-padding before and after63    pad = [[float('inf')]*3, [float('inf')]*3, [float('inf')]*3]64    coords.insert(0, pad)65    coords.append(pad)66 67    if target == "ParD2":68        coords = [pad, pad] + coords + [pad, pad]69    elif target == "ParD3":70        coords = [pad]*2 + coords + [pad]*671    elif target == "TrpB4":72        coords = [pad] + coords73 74    coords_tensor = torch.tensor(coords, device=device).unsqueeze(0)  # (1, N, 3, 3)75 76    return coords_tensor, residue_types77 78num_replicates = 1079campaign_number = 2 # change this according to the campaign we are interested in80dataset_size = 96 # change this according to the dataset size we are interested in81 82sequence_tokenizer = EsmSequenceTokenizer()83 84datasets = ["GB1", "TrpB4"]85data_root_path = "/global/cfs/projectdirs/m4235/sebastian/data"86 87for data in datasets:88    print(data)89    for i in range(num_replicates):90        cfg_filename = f"./config.yaml"91        cfg = OmegaConf.load(cfg_filename)92        sampling_temperature=193        OmegaConf.update(cfg, "train.lightning_model_args.sampling_temperature", sampling_temperature)94        mask_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["mask"]95        bos_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["bos"]96        eos_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["eos"]97        pad_token_sequence = cfg["nn"]["model_args"]["residue_token_info"]["pad"]98        99        if not data.startswith("TrpB"):100            df = pd.read_csv(f"{data_root_path}/{data}/scale2max/{data}.csv")101            with open(f"{data_root_path}/{data}/{data}.fasta", "r") as file:102                parent_sequence_decoded = file.readlines()[1].strip()103                104        else:105            df = pd.read_csv(f"{data_root_path}/TrpB/scale2max/{data}.csv")106            with open(f"{data_root_path}/TrpB/TrpB.fasta", "r") as file:107                parent_sequence_decoded = file.readlines()[1].strip()108                109        if data != "GB1":        110            muts = df["muts"].iloc[0]111        else:112            muts = df["muts"].iloc[100000]113        114        numbers = re.findall(r'\d+', muts)115        mask_indices = list(map(int, numbers))116        # mask_indices = [i-1 for i in mask_indices] #convert to 0-based indexing117        118        fitness_scores = []119        120# Load from base_model_{dataset_size}121        trpb_base = torch.load(f"./{data}/base_model_{dataset_size}/trpb_post_rd_{campaign_number-1}_{i}.pt")122        all_unmasked_sequences_decoded_base = trpb_base["all_unmasked_sequences_decoded"]123        all_unmasked_sequences_base = trpb_base["all_unmasked_sequences"]124        all_masked_sequences_base = trpb_base["all_masked_sequences"]125        all_unmasked_sequences_base = all_unmasked_sequences_base.reshape(-1, all_unmasked_sequences_base.shape[-1])126        all_logps_base = trpb_base["all_logps"]127        128        for unmasked_sequence_decoded, unmasked_sequence in zip(all_unmasked_sequences_decoded_base, all_unmasked_sequences_base):129            index_residue_0 = unmasked_sequence_decoded[mask_indices[0]-1]130            index_residue_1 = unmasked_sequence_decoded[mask_indices[1]-1]131            index_residue_2 = unmasked_sequence_decoded[mask_indices[2]-1]132            try:133                index_residue_3 = unmasked_sequence_decoded[mask_indices[3]-1]134                mutations = [index_residue_0, index_residue_1, index_residue_2, index_residue_3]135                muts = ''.join(mutations)136            except:137                mutations = [index_residue_0, index_residue_1, index_residue_2]138                muts = ''.join(mutations)139            140            df_filtered = df[df["AAs"] == muts]141            142            if len(df_filtered) == 0:143                if torch.any((unmasked_sequence[1:-1] > 23) | (unmasked_sequence[1:-1] < 4)):144                    print(f"Invalid sequence {muts}")145                    fitness_score = -2146                else:147                    print(f"Invalid sequence {muts}")148                    fitness_score = -2149            else:150                fitness_score = df_filtered["fitness"].values[0]151            fitness_scores.append(fitness_score)152        153        # Load from aligned_0_{dataset_size}154        trpb_aligned = torch.load(f"./{data}/aligned_{campaign_number-2}_{dataset_size}_{i}/trpb_post_rd_{campaign_number-1}_{i}.pt")155        all_unmasked_sequences_decoded_aligned_0 = trpb_aligned["all_unmasked_sequences_decoded"]156        all_unmasked_sequences_aligned_0 = trpb_aligned["all_unmasked_sequences"]157        all_masked_sequences_aligned_0 = trpb_aligned["all_masked_sequences"]158        all_unmasked_sequences_aligned_0 = all_unmasked_sequences_aligned_0.reshape(-1, all_unmasked_sequences_aligned_0.shape[-1])159        all_logps_aligned_0 = trpb_aligned["all_logps"]160        161        for unmasked_sequence_decoded, unmasked_sequence in zip(all_unmasked_sequences_decoded_aligned_0, all_unmasked_sequences_aligned_0):162            index_residue_0 = unmasked_sequence_decoded[mask_indices[0]-1]163            index_residue_1 = unmasked_sequence_decoded[mask_indices[1]-1]164            index_residue_2 = unmasked_sequence_decoded[mask_indices[2]-1]165            try:166                index_residue_3 = unmasked_sequence_decoded[mask_indices[3]-1]167                mutations = [index_residue_0, index_residue_1, index_residue_2, index_residue_3]168                muts = ''.join(mutations)169            except:170                mutations = [index_residue_0, index_residue_1, index_residue_2]171                muts = ''.join(mutations)172            173            df_filtered = df[df["AAs"] == muts]174            175            if len(df_filtered) == 0:176                if torch.any((unmasked_sequence[1:-1] > 23) | (unmasked_sequence[1:-1] < 4)):177                    print(f"Invalid sequence {muts}")178                    fitness_score = -2179                else:180                    print(f"Invalid sequence {muts}")181                    fitness_score = -2182            else:183                fitness_score = df_filtered["fitness"].values[0]184            fitness_scores.append(fitness_score)185        186        # Load from aligned_1_{dataset_size}187        trpb_aligned_1 = torch.load(f"./{data}/aligned_{campaign_number-1}_{dataset_size}_{i}/trpb_{i}.pt")188        all_unmasked_sequences_decoded_aligned_1 = trpb_aligned_1["all_unmasked_sequences_decoded"]189        all_unmasked_sequences_aligned_1 = trpb_aligned_1["all_unmasked_sequences"]190        all_masked_sequences_aligned_1 = trpb_aligned_1["all_masked_sequences"]191        all_unmasked_sequences_aligned_1 = all_unmasked_sequences_aligned_1.reshape(-1, all_unmasked_sequences_aligned_1.shape[-1])192        all_logps_aligned_1 = trpb_aligned_1["all_logps"]193 194        for unmasked_sequence_decoded, unmasked_sequence in zip(all_unmasked_sequences_decoded_aligned_1, all_unmasked_sequences_aligned_1):195            index_residue_0 = unmasked_sequence_decoded[mask_indices[0]-1]196            index_residue_1 = unmasked_sequence_decoded[mask_indices[1]-1]197            index_residue_2 = unmasked_sequence_decoded[mask_indices[2]-1]198            try:199                index_residue_3 = unmasked_sequence_decoded[mask_indices[3]-1]200                mutations = [index_residue_0, index_residue_1, index_residue_2, index_residue_3]201                muts = ''.join(mutations)202            except:203                mutations = [index_residue_0, index_residue_1, index_residue_2]204                muts = ''.join(mutations)205 206            df_filtered = df[df["AAs"] == muts]207 208            if len(df_filtered) == 0:209                if torch.any((unmasked_sequence[1:-1] > 23) | (unmasked_sequence[1:-1] < 4)):210                    print(f"Invalid sequence {muts}")211                    fitness_score = -2212                else:213                    print(f"Invalid sequence {muts}")214                    fitness_score = -2215            else:216                fitness_score = df_filtered["fitness"].values[0]217            fitness_scores.append(fitness_score)218            219        # Concatenate the sequences and logps from all models220        all_unmasked_sequences = torch.cat((all_unmasked_sequences_base, all_unmasked_sequences_aligned_0, all_unmasked_sequences_aligned_1),dim=0)221        all_masked_sequences = torch.cat((all_masked_sequences_base, all_masked_sequences_aligned_0, all_masked_sequences_aligned_1),dim=0)222        print(all_logps_base.shape, all_logps_aligned_0.shape, all_logps_aligned_1.shape)223        all_logps = torch.cat((all_logps_base, all_logps_aligned_0, all_logps_aligned_1),dim=0)224 225        all_fitness_scores = fitness_scores226        227        # Check for duplicates in all_unmasked_sequences228        unique_sequences, counts = torch.unique(all_unmasked_sequences, dim=0, return_counts=True)229        num_duplicates = torch.sum(counts > 1).item()230        print(f"Number of duplicate sequences: {num_duplicates}")231        232        all_fitness_scores = np.array(all_fitness_scores)233        all_fitness_scores = np.where(all_fitness_scores > 0, -np.log(all_fitness_scores), 10)234        235        sampling_temperature = 1 # hard-coding a sampling temperature of 1 for mixed-temperature alignment236 237        sequence_length = all_unmasked_sequences.shape[1]238 239        sequence_id = torch.ones((all_unmasked_sequences.shape[0], sequence_length), device=device).long() * 1240 241        structure_tokens = torch.ones((1, sequence_length), device=device).long() * 4096242        structure_tokens[:, 0] = 4098243        structure_tokens[:, -1] = 4097244 245        coords, residue_types = get_backbone_coords_from_local_pdb(f"{data_root_path}/{data}/{data}.pdb", chain_id='A', sequence_length=sequence_length-2, target=data) if not data.startswith("TrpB") else get_backbone_coords_from_local_pdb(f"{data_root_path}/TrpB/TrpB.pdb", chain_id='A', sequence_length=sequence_length-2, target=data)246 247        # parent sequence sanity check248        coords_trimmed = coords[:, 1:-1]  # shape: (1, N-2, 3, 3)249 250        # Step 2: Determine mask of non-padding residues (i.e., not all coords are inf)251        valid_mask = ~(torch.isinf(coords_trimmed).view(-1, 9).any(dim=1))  # shape: (N-2,)252        residues_to_compare = [r for r, valid in zip(list(parent_sequence_decoded), valid_mask) if valid]253 254        if residue_types != residues_to_compare:255            print("Residue mismatch detected!")256            for i, (ref, pdb) in enumerate(zip(residues_to_compare, residue_types)):257                if ref != pdb:258                    print(f"Position {i}: expected {ref}, got {pdb}")259        else:260            print("Residues match.")261            print(coords.shape)262 263        assert coords.shape[1] == sequence_length, f"Coords length {coords.shape[1]} does not match sequence length {sequence_length}"264 265        average_plddt = torch.ones((1), device=device)266 267        per_res_plddt = torch.zeros((1, sequence_length), device=device)268        ss8_tokens = torch.zeros((1, sequence_length), device=device).long()269        sasa_tokens = torch.zeros((1, sequence_length), device=device).long()270 271        function_tokens = torch.zeros((1, sequence_length, 8), device=device).long()272        residue_annotation_tokens = torch.zeros((1, sequence_length, 16), device=device).long()273                274        275        with h5py.File(f"./{data}/alignment_dataset_{campaign_number}_{dataset_size}_from_ESM3_{i}.hdf5", "w") as f:276            masked_sequence_tokens = f.create_dataset("masked_sequence_tokens", data=all_masked_sequences.cpu().numpy())277            unmasked_sequence_tokens = f.create_dataset("unmasked_sequence_tokens", data=all_unmasked_sequences.cpu().numpy())278            sequence_id = f.create_dataset("sequence_id", data=sequence_id.cpu().numpy())279            structure_tokens = f.create_dataset("structural_tokens", data=structure_tokens.cpu().numpy())280            coords = f.create_dataset("bb_coords", data=coords.cpu().numpy())281            average_plddt = f.create_dataset("average_plddt", data=average_plddt.cpu().numpy())282            per_res_plddt = f.create_dataset("per_res_plddt", data=per_res_plddt.cpu().numpy())283            ss8_tokens = f.create_dataset("ss8_tokens", data=ss8_tokens.cpu().numpy())284            sasa_tokens = f.create_dataset("sasa_tokens", data=sasa_tokens.cpu().numpy())285            function_tokens = f.create_dataset("function_tokens", data=function_tokens.cpu().numpy())286            residue_annotation_tokens = f.create_dataset("residue_annotation_tokens", data=residue_annotation_tokens.cpu().numpy())287 288            ref_logps = f.create_dataset("ref_logps", data=all_logps.cpu().numpy())289            energies = f.create_dataset("energies", data=all_fitness_scores)290 291 292            f.attrs["num_prompts"] = 1293            f.attrs["num_examples_per_prompt"] = masked_sequence_tokens.shape[0]294            f.attrs["fixed_bb_coords"] = True295            f.attrs["fixed_average_plddt"] = True296            f.attrs["fixed_per_res_plddt"] = True297            f.attrs["fixed_ss8_tokens"] = True298            f.attrs["fixed_sasa_tokens"] = True299            f.attrs["fixed_function_tokens"] = True300            f.attrs["fixed_residue_annotation_tokens"] = True301            f.attrs["fixed_structural_tokens"] = True302            f.attrs["sampling_temperature"] = sampling_temperature303    304