Team Ai
Modelpublic

OneScience-Group/CodonTransformer

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes7downloads
CodonPrediction.py856 linesDownload Raw Back to CodonTransformer
1"""2File: CodonPrediction.py3---------------------------4Includes functions to tokenize input, load models, infer predicted dna sequences and5helper functions related to processing data for passing to the model.6"""7 8import warnings9from typing import Any, Dict, List, Optional, Tuple, Union10 11import numpy as np12import onnxruntime as rt13import torch14import transformers15from transformers import (16    AutoTokenizer,17    BatchEncoding,18    BigBirdConfig,19    BigBirdForMaskedLM,20    PreTrainedTokenizerFast,21)22 23from CodonTransformer.CodonData import get_merged_seq24from CodonTransformer.CodonUtils import (25    AMINO_ACID_TO_INDEX,26    INDEX2TOKEN,27    NUM_ORGANISMS,28    ORGANISM2ID,29    TOKEN2INDEX,30    DNASequencePrediction,31)32 33 34def predict_dna_sequence(35    protein: str,36    organism: Union[int, str],37    device: torch.device,38    tokenizer: Union[str, PreTrainedTokenizerFast] = None,39    model: Union[str, torch.nn.Module] = None,40    attention_type: str = "original_full",41    deterministic: bool = True,42    temperature: float = 0.2,43    top_p: float = 0.95,44    num_sequences: int = 1,45    match_protein: bool = False,46) -> Union[DNASequencePrediction, List[DNASequencePrediction]]:47    """48    Predict the DNA sequence(s) for a given protein using the CodonTransformer model.49 50    This function takes a protein sequence and an organism (as ID or name) as input51    and returns the predicted DNA sequence(s) using the CodonTransformer model. It can use52    either provided tokenizer and model objects or load them from specified paths.53 54    Args:55        protein (str): The input protein sequence for which to predict the DNA sequence.56        organism (Union[int, str]): Either the ID of the organism or its name (e.g.,57            "Escherichia coli general"). If a string is provided, it will be converted58            to the corresponding ID using ORGANISM2ID.59        device (torch.device): The device (CPU or GPU) to run the model on.60        tokenizer (Union[str, PreTrainedTokenizerFast, None], optional): Either a file61            path to load the tokenizer from, a pre-loaded tokenizer object, or None. If62            None, it will be loaded from HuggingFace. Defaults to None.63        model (Union[str, torch.nn.Module, None], optional): Either a file path to load64            the model from, a pre-loaded model object, or None. If None, it will be65            loaded from HuggingFace. Defaults to None.66        attention_type (str, optional): The type of attention mechanism to use in the67            model. Can be either 'block_sparse' or 'original_full'. Defaults to68            "original_full".69        deterministic (bool, optional): Whether to use deterministic decoding (most70            likely tokens). If False, samples tokens according to their probabilities71            adjusted by the temperature. Defaults to True.72        temperature (float, optional): A value controlling the randomness of predictions73            during non-deterministic decoding. Lower values (e.g., 0.2) make the model74            more conservative, while higher values (e.g., 0.8) increase randomness.75            Using high temperatures may result in prediction of DNA sequences that76            do not translate to the input protein.77            Recommended values are:78                - Low randomness: 0.279                - Medium randomness: 0.580                - High randomness: 0.881            The temperature must be a positive float. Defaults to 0.2.82        top_p (float, optional): The cumulative probability threshold for nucleus sampling.83            Tokens with cumulative probability up to top_p are considered for sampling.84            This parameter helps balance diversity and coherence in the predicted DNA sequences.85            The value must be a float between 0 and 1. Defaults to 0.95.86        num_sequences (int, optional): The number of DNA sequences to generate. Only applicable87            when deterministic is False. Defaults to 1.88        match_protein (bool, optional): Ensures the predicted DNA sequence is translated89            to the input protein sequence by sampling from only the respective codons of90            given amino acids. Defaults to False.91 92    Returns:93        Union[DNASequencePrediction, List[DNASequencePrediction]]: An object or list of objects94        containing the prediction results:95            - organism (str): Name of the organism used for prediction.96            - protein (str): Input protein sequence for which DNA sequence is predicted.97            - processed_input (str): Processed input sequence (merged protein and DNA).98            - predicted_dna (str): Predicted DNA sequence.99 100    Raises:101        ValueError: If the protein sequence is empty, if the organism is invalid,102            if the temperature is not a positive float, if top_p is not between 0 and 1,103            or if num_sequences is less than 1 or used with deterministic mode.104 105    Note:106        This function uses ORGANISM2ID, INDEX2TOKEN, and AMINO_ACID_TO_INDEX dictionaries107        imported from CodonTransformer.CodonUtils. ORGANISM2ID maps organism names to their108        corresponding IDs. INDEX2TOKEN maps model output indices (token IDs) to109        respective codons. AMINO_ACID_TO_INDEX maps each amino acid and stop symbol to indices110        of codon tokens that translate to it.111 112    Example:113        >>> import torch114        >>> from transformers import AutoTokenizer, BigBirdForMaskedLM115        >>> from CodonTransformer.CodonPrediction import predict_dna_sequence116        >>> from CodonTransformer.CodonJupyter import format_model_output117        >>>118        >>> # Set up device119        >>> device = torch.device("cuda" if torch.cuda.is_available() else "cpu")120        >>>121        >>> # Load tokenizer and model122        >>> tokenizer = AutoTokenizer.from_pretrained("adibvafa/CodonTransformer")123        >>> model = BigBirdForMaskedLM.from_pretrained("adibvafa/CodonTransformer")124        >>> model = model.to(device)125        >>>126        >>> # Define protein sequence and organism127        >>> protein = "MKTVRQERLKSIVRILERSKEPVSGAQLAEELSVSRQVIVQDIAYLRSLGYNIVATPRGYVLA"128        >>> organism = "Escherichia coli general"129        >>>130        >>> # Predict DNA sequence with deterministic decoding (single sequence)131        >>> output = predict_dna_sequence(132        ...     protein=protein,133        ...     organism=organism,134        ...     device=device,135        ...     tokenizer=tokenizer,136        ...     model=model,137        ...     attention_type="original_full",138        ...     deterministic=True139        ... )140        >>>141        >>> # Predict multiple DNA sequences with low randomness and top_p sampling142        >>> output_random = predict_dna_sequence(143        ...     protein=protein,144        ...     organism=organism,145        ...     device=device,146        ...     tokenizer=tokenizer,147        ...     model=model,148        ...     attention_type="original_full",149        ...     deterministic=False,150        ...     temperature=0.2,151        ...     top_p=0.95,152        ...     num_sequences=3153        ... )154        >>>155        >>> print(format_model_output(output))156        >>> for i, seq in enumerate(output_random, 1):157        ...     print(f"Sequence {i}:")158        ...     print(format_model_output(seq))159        ...     print()160    """161    if not protein:162        raise ValueError("Protein sequence cannot be empty.")163 164    if not isinstance(temperature, (float, int)) or temperature <= 0:165        raise ValueError("Temperature must be a positive float.")166 167    if not isinstance(top_p, (float, int)) or not 0 < top_p <= 1.0:168        raise ValueError("top_p must be a float between 0 and 1.")169 170    if not isinstance(num_sequences, int) or num_sequences < 1:171        raise ValueError("num_sequences must be a positive integer.")172 173    if deterministic and num_sequences > 1:174        raise ValueError(175            "Multiple sequences can only be generated in non-deterministic mode."176        )177 178    # Load tokenizer179    if not isinstance(tokenizer, PreTrainedTokenizerFast):180        tokenizer = load_tokenizer(tokenizer)181 182    # Load model183    if not isinstance(model, torch.nn.Module):184        model = load_model(model, device=device, attention_type=attention_type)185    else:186        model.eval()187        model.bert.set_attention_type(attention_type)188        model.to(device)189 190    # Validate organism and convert to organism_id and organism_name191    organism_id, organism_name = validate_and_convert_organism(organism)192 193    # Inference loop194    with torch.no_grad():195        # Tokenize the input sequence196        merged_seq = get_merged_seq(protein=protein, dna="")197        input_dict = {198            "idx": 0,  # sample index199            "codons": merged_seq,200            "organism": organism_id,201        }202        tokenized_input = tokenize([input_dict], tokenizer=tokenizer).to(device)203 204        # Get the model predictions205        output_dict = model(**tokenized_input, return_dict=True)206        logits = output_dict.logits.detach().cpu()207        logits = logits[:, 1:-1, :]  # Remove [CLS] and [SEP] tokens208 209        # Mask the logits of codons that do not correspond to the input protein sequence210        if match_protein:211            possible_tokens_per_position = [212                AMINO_ACID_TO_INDEX[token[0]] for token in merged_seq.split(" ")213            ]214            mask = torch.full_like(logits, float("-inf"))215 216            for pos, possible_tokens in enumerate(possible_tokens_per_position):217                mask[:, pos, possible_tokens] = 0218 219            logits = mask + logits220 221        predictions = []222        for _ in range(num_sequences):223            # Decode the predicted DNA sequence from the model output224            if deterministic:225                predicted_indices = logits.argmax(dim=-1).squeeze().tolist()226            else:227                predicted_indices = sample_non_deterministic(228                    logits=logits, temperature=temperature, top_p=top_p229                )230 231            predicted_dna = list(map(INDEX2TOKEN.__getitem__, predicted_indices))232            predicted_dna = (233                "".join([token[-3:] for token in predicted_dna]).strip().upper()234            )235 236            predictions.append(237                DNASequencePrediction(238                    organism=organism_name,239                    protein=protein,240                    processed_input=merged_seq,241                    predicted_dna=predicted_dna,242                )243            )244 245    return predictions[0] if num_sequences == 1 else predictions246 247 248def sample_non_deterministic(249    logits: torch.Tensor,250    temperature: float = 0.2,251    top_p: float = 0.95,252) -> List[int]:253    """254    Sample token indices from logits using temperature scaling and nucleus (top-p) sampling.255 256    This function applies temperature scaling to the logits, computes probabilities,257    and then performs nucleus sampling to select token indices. It is used for258    non-deterministic decoding in language models to introduce randomness while259    maintaining coherence in the generated sequences.260 261    Args:262        logits (torch.Tensor): The logits output from the model of shape263            [seq_len, vocab_size] or [batch_size, seq_len, vocab_size].264        temperature (float, optional): Temperature value for scaling logits.265            Must be a positive float. Defaults to 1.0.266        top_p (float, optional): Cumulative probability threshold for nucleus sampling.267            Must be a float between 0 and 1. Tokens with cumulative probability up to268            `top_p` are considered for sampling. Defaults to 0.95.269 270    Returns:271        List[int]: A list of sampled token indices corresponding to the predicted tokens.272 273    Raises:274        ValueError: If `temperature` is not a positive float or if `top_p` is not between 0 and 1.275 276    Example:277        >>> logits = model_output.logits  # Assume logits is a tensor of shape [seq_len, vocab_size]278        >>> predicted_indices = sample_non_deterministic(logits, temperature=0.7, top_p=0.9)279    """280    if not isinstance(temperature, (float, int)) or temperature <= 0:281        raise ValueError("Temperature must be a positive float.")282 283    if not isinstance(top_p, (float, int)) or not 0 < top_p <= 1.0:284        raise ValueError("top_p must be a float between 0 and 1.")285 286    # Compute probabilities using temperature scaling287    probs = torch.softmax(logits / temperature, dim=-1)288 289 290    # Remove batch dimension if present291    if probs.dim() == 3:292        probs = probs.squeeze(0)  # Shape: [seq_len, vocab_size]293 294    # Sort probabilities in descending order295    probs_sort, probs_idx = torch.sort(probs, dim=-1, descending=True)296    probs_sum = torch.cumsum(probs_sort, dim=-1)297    mask = probs_sum - probs_sort > top_p298 299    # Zero out probabilities for tokens beyond the top-p threshold300    probs_sort[mask] = 0.0301 302    # Renormalize the probabilities303    probs_sort.div_(probs_sort.sum(dim=-1, keepdim=True))304    next_token = torch.multinomial(probs_sort, num_samples=1)305    predicted_indices = torch.gather(probs_idx, -1, next_token).squeeze(-1)306 307    return predicted_indices.tolist()308 309 310def load_model(311    model_path: Optional[str] = None,312    device: torch.device = None,313    attention_type: str = "original_full",314    num_organisms: int = None,315    remove_prefix: bool = True,316) -> torch.nn.Module:317    """318    Load a BigBirdForMaskedLM model from a model file, checkpoint, or HuggingFace.319 320    Args:321        model_path (Optional[str]): Path to the model file or checkpoint. If None,322            load from HuggingFace.323        device (torch.device, optional): The device to load the model onto.324        attention_type (str, optional): The type of attention, 'block_sparse'325            or 'original_full'.326        num_organisms (int, optional): Number of organisms, needed if loading from a327            checkpoint that requires this.328        remove_prefix (bool, optional): Whether to remove the "model." prefix from the329            keys in the state dict.330 331    Returns:332        torch.nn.Module: The loaded model.333    """334    if not model_path:335        warnings.warn("Model path not provided. Loading from HuggingFace.", UserWarning)336        model = BigBirdForMaskedLM.from_pretrained("adibvafa/CodonTransformer")337 338    elif model_path.endswith(".ckpt"):339        checkpoint = torch.load(model_path)340        state_dict = checkpoint["state_dict"]341 342        # Remove the "model." prefix from the keys343        if remove_prefix:344            state_dict = {345                key.replace("model.", ""): value for key, value in state_dict.items()346            }347 348        if num_organisms is None:349            num_organisms = NUM_ORGANISMS350 351        # Load model configuration and instantiate the model352        config = load_bigbird_config(num_organisms)353        model = BigBirdForMaskedLM(config=config)354        model.load_state_dict(state_dict)355 356    elif model_path.endswith(".pt"):357        state_dict = torch.load(model_path)358        config = state_dict.pop("self.config")359        model = BigBirdForMaskedLM(config=config)360        model.load_state_dict(state_dict)361 362    else:363        raise ValueError(364            "Unsupported file type. Please provide a .ckpt or .pt file, "365            "or None to load from HuggingFace."366        )367 368    # Prepare model for evaluation369    model.bert.set_attention_type(attention_type)370    model.eval()371    if device:372        model.to(device)373 374    return model375 376 377def load_bigbird_config(num_organisms: int) -> BigBirdConfig:378    """379    Load the config object used to train the BigBird transformer.380 381    Args:382        num_organisms (int): The number of organisms.383 384    Returns:385        BigBirdConfig: The configuration object for BigBird.386    """387    config = transformers.BigBirdConfig(388        vocab_size=len(TOKEN2INDEX),  # Equal to len(tokenizer)389        type_vocab_size=num_organisms,390        sep_token_id=2,391    )392    return config393 394 395def create_model_from_checkpoint(396    checkpoint_dir: str, output_model_dir: str, num_organisms: int397) -> None:398    """399    Save a model to disk using a previous checkpoint.400 401    Args:402        checkpoint_dir (str): Directory where the checkpoint is stored.403        output_model_dir (str): Directory where the model will be saved.404        num_organisms (int): Number of organisms.405    """406    checkpoint = load_model(model_path=checkpoint_dir, num_organisms=num_organisms)407    state_dict = checkpoint.state_dict()408    state_dict["self.config"] = load_bigbird_config(num_organisms=num_organisms)409 410    # Save the model state dict to the output directory411    torch.save(state_dict, output_model_dir)412 413 414def load_tokenizer(tokenizer_path: Optional[str] = None) -> PreTrainedTokenizerFast:415    """416    Create and return a tokenizer object from tokenizer path or HuggingFace.417 418    Args:419        tokenizer_path (Optional[str]): Path to the tokenizer file. If None,420        load from HuggingFace.421 422    Returns:423        PreTrainedTokenizerFast: The tokenizer object.424    """425    if not tokenizer_path:426        warnings.warn(427            "Tokenizer path not provided. Loading from HuggingFace.", UserWarning428        )429        return AutoTokenizer.from_pretrained("adibvafa/CodonTransformer")430 431    return transformers.PreTrainedTokenizerFast(432        tokenizer_file=tokenizer_path,433        bos_token="[CLS]",434        eos_token="[SEP]",435        unk_token="[UNK]",436        sep_token="[SEP]",437        pad_token="[PAD]",438        cls_token="[CLS]",439        mask_token="[MASK]",440    )441 442 443def tokenize(444    batch: List[Dict[str, Any]],445    tokenizer: Union[PreTrainedTokenizerFast, str] = None,446    max_len: int = 2048,447) -> BatchEncoding:448    """449    Return the tokenized sequences given a batch of input data.450    Each data in the batch is expected to be a dictionary with "codons" and451    "organism" keys.452 453    Args:454        batch (List[Dict[str, Any]]): A list of dictionaries with "codons" and455            "organism" keys.456        tokenizer (PreTrainedTokenizerFast, str, optional): The tokenizer object or457            path to the tokenizer file.458        max_len (int, optional): Maximum length of the tokenized sequence.459 460    Returns:461        BatchEncoding: The tokenized batch.462    """463    if not isinstance(tokenizer, PreTrainedTokenizerFast):464        tokenizer = load_tokenizer(tokenizer)465 466    tokenized = tokenizer(467        [data["codons"] for data in batch],468        return_attention_mask=True,469        return_token_type_ids=True,470        truncation=True,471        padding=True,472        max_length=max_len,473        return_tensors="pt",474    )475 476    # Add token type IDs for species477    seq_len = tokenized["input_ids"].shape[-1]478    species_index = torch.tensor([[data["organism"]] for data in batch])479    tokenized["token_type_ids"] = species_index.repeat(1, seq_len)480 481    return tokenized482 483 484def validate_and_convert_organism(organism: Union[int, str]) -> Tuple[int, str]:485    """486    Validate and convert the organism input to both ID and name.487 488    This function takes either an organism ID or name as input and returns both489    the ID and name. It performs validation to ensure the input corresponds to490    a valid organism in the ORGANISM2ID dictionary.491 492    Args:493        organism (Union[int, str]): Either the ID of the organism (int) or its494        name (str).495 496    Returns:497        Tuple[int, str]: A tuple containing the organism ID (int) and name (str).498 499    Raises:500        ValueError: If the input is neither a string nor an integer, if the501        organism name is not found in ORGANISM2ID, if the organism ID is not a502        value in ORGANISM2ID, or if no name is found for a given ID.503 504    Note:505        This function relies on the ORGANISM2ID dictionary imported from506        CodonTransformer.CodonUtils, which maps organism names to their507        corresponding IDs.508    """509    if isinstance(organism, str):510        if organism not in ORGANISM2ID:511            raise ValueError(512                f"Invalid organism name: {organism}. "513                "Please use a valid organism name or ID."514            )515        organism_id = ORGANISM2ID[organism]516        organism_name = organism517 518    elif isinstance(organism, int):519        if organism not in ORGANISM2ID.values():520            raise ValueError(521                f"Invalid organism ID: {organism}. "522                "Please use a valid organism name or ID."523            )524 525        organism_id = organism526        organism_name = next(527            (name for name, id in ORGANISM2ID.items() if id == organism), None528        )529        if organism_name is None:530            raise ValueError(f"No organism name found for ID: {organism}")531 532    return organism_id, organism_name533 534 535def get_high_frequency_choice_sequence(536    protein: str, codon_frequencies: Dict[str, Tuple[List[str], List[float]]]537) -> str:538    """539    Return the DNA sequence optimized using High Frequency Choice (HFC) approach540    in which the most frequent codon for a given amino acid is always chosen.541 542    Args:543        protein (str): The protein sequence.544        codon_frequencies (Dict[str, Tuple[List[str], List[float]]]): Codon545        frequencies for each amino acid.546 547    Returns:548        str: The optimized DNA sequence.549    """550    # Select the most frequent codon for each amino acid in the protein sequence551    dna_codons = [552        codon_frequencies[aminoacid][0][np.argmax(codon_frequencies[aminoacid][1])]553        for aminoacid in protein554    ]555    return "".join(dna_codons)556 557 558def precompute_most_frequent_codons(559    codon_frequencies: Dict[str, Tuple[List[str], List[float]]],560) -> Dict[str, str]:561    """562    Precompute the most frequent codon for each amino acid.563 564    Args:565        codon_frequencies (Dict[str, Tuple[List[str], List[float]]]): Codon566        frequencies for each amino acid.567 568    Returns:569        Dict[str, str]: The most frequent codon for each amino acid.570    """571    # Create a dictionary mapping each amino acid to its most frequent codon572    return {573        aminoacid: codons[np.argmax(frequencies)]574        for aminoacid, (codons, frequencies) in codon_frequencies.items()575    }576 577 578def get_high_frequency_choice_sequence_optimized(579    protein: str, codon_frequencies: Dict[str, Tuple[List[str], List[float]]]580) -> str:581    """582    Efficient implementation of get_high_frequency_choice_sequence that uses583    vectorized operations and helper functions, achieving up to x10 faster speed.584 585    Args:586        protein (str): The protein sequence.587        codon_frequencies (Dict[str, Tuple[List[str], List[float]]]): Codon588        frequencies for each amino acid.589 590    Returns:591        str: The optimized DNA sequence.592    """593    # Precompute the most frequent codons for each amino acid594    most_frequent_codons = precompute_most_frequent_codons(codon_frequencies)595 596    return "".join(most_frequent_codons[aminoacid] for aminoacid in protein)597 598 599def get_background_frequency_choice_sequence(600    protein: str, codon_frequencies: Dict[str, Tuple[List[str], List[float]]]601) -> str:602    """603    Return the DNA sequence optimized using Background Frequency Choice (BFC)604    approach in which a random codon for a given amino acid is chosen using605    the codon frequencies probability distribution.606 607    Args:608        protein (str): The protein sequence.609        codon_frequencies (Dict[str, Tuple[List[str], List[float]]]): Codon610        frequencies for each amino acid.611 612    Returns:613        str: The optimized DNA sequence.614    """615    # Select a random codon for each amino acid based on the codon frequencies616    # probability distribution617    dna_codons = [618        np.random.choice(619            codon_frequencies[aminoacid][0], p=codon_frequencies[aminoacid][1]620        )621        for aminoacid in protein622    ]623    return "".join(dna_codons)624 625 626def precompute_cdf(627    codon_frequencies: Dict[str, Tuple[List[str], List[float]]],628) -> Dict[str, Tuple[List[str], Any]]:629    """630    Precompute the cumulative distribution function (CDF) for each amino acid.631 632    Args:633        codon_frequencies (Dict[str, Tuple[List[str], List[float]]]): Codon634        frequencies for each amino acid.635 636    Returns:637        Dict[str, Tuple[List[str], Any]]: CDFs for each amino acid.638    """639    cdf = {}640 641    # Calculate the cumulative distribution function for each amino acid642    for aminoacid, (codons, frequencies) in codon_frequencies.items():643        cdf[aminoacid] = (codons, np.cumsum(frequencies))644 645    return cdf646 647 648def get_background_frequency_choice_sequence_optimized(649    protein: str, codon_frequencies: Dict[str, Tuple[List[str], List[float]]]650) -> str:651    """652    Efficient implementation of get_background_frequency_choice_sequence that uses653    vectorized operations and helper functions, achieving up to x8 faster speed.654 655    Args:656        protein (str): The protein sequence.657        codon_frequencies (Dict[str, Tuple[List[str], List[float]]]): Codon658        frequencies for each amino acid.659 660    Returns:661        str: The optimized DNA sequence.662    """663    dna_codons = []664    cdf = precompute_cdf(codon_frequencies)665 666    # Select a random codon for each amino acid using the precomputed CDFs667    for aminoacid in protein:668        codons, cumulative_prob = cdf[aminoacid]669        selected_codon_index = np.searchsorted(cumulative_prob, np.random.rand())670        dna_codons.append(codons[selected_codon_index])671 672    return "".join(dna_codons)673 674 675def get_uniform_random_choice_sequence(676    protein: str, codon_frequencies: Dict[str, Tuple[List[str], List[float]]]677) -> str:678    """679    Return the DNA sequence optimized using Uniform Random Choice (URC) approach680    in which a random codon for a given amino acid is chosen using a uniform681    prior.682 683    Args:684        protein (str): The protein sequence.685        codon_frequencies (Dict[str, Tuple[List[str], List[float]]]): Codon686        frequencies for each amino acid.687 688    Returns:689        str: The optimized DNA sequence.690    """691    # Select a random codon for each amino acid using a uniform prior distribution692    dna_codons = []693    for aminoacid in protein:694        codons = codon_frequencies[aminoacid][0]695        random_index = np.random.randint(0, len(codons))696        dna_codons.append(codons[random_index])697    return "".join(dna_codons)698 699 700def get_icor_prediction(input_seq: str, model_path: str, stop_symbol: str) -> str:701    """702    Return the optimized codon sequence for the given protein sequence using ICOR.703 704    Credit: ICOR: improving codon optimization with recurrent neural networks705            Rishab Jain, Aditya Jain, Elizabeth Mauro, Kevin LeShane, Douglas706            Densmore707 708    Args:709        input_seq (str): The input protein sequence.710        model_path (str): The path to the ICOR model.711        stop_symbol (str): The symbol representing stop codons in the sequence.712 713    Returns:714        str: The optimized DNA sequence.715    """716    input_seq = input_seq.strip().upper()717    input_seq = input_seq.replace(stop_symbol, "*")718 719    # Define categorical labels from when model was trained.720    labels = [721        "AAA",722        "AAC",723        "AAG",724        "AAT",725        "ACA",726        "ACG",727        "ACT",728        "AGC",729        "ATA",730        "ATC",731        "ATG",732        "ATT",733        "CAA",734        "CAC",735        "CAG",736        "CCG",737        "CCT",738        "CTA",739        "CTC",740        "CTG",741        "CTT",742        "GAA",743        "GAT",744        "GCA",745        "GCC",746        "GCG",747        "GCT",748        "GGA",749        "GGC",750        "GTC",751        "GTG",752        "GTT",753        "TAA",754        "TAT",755        "TCA",756        "TCG",757        "TCT",758        "TGG",759        "TGT",760        "TTA",761        "TTC",762        "TTG",763        "TTT",764        "ACC",765        "CAT",766        "CCA",767        "CGG",768        "CGT",769        "GAC",770        "GAG",771        "GGT",772        "AGT",773        "GGG",774        "GTA",775        "TGC",776        "CCC",777        "CGA",778        "CGC",779        "TAC",780        "TAG",781        "TCC",782        "AGA",783        "AGG",784        "TGA",785    ]786 787    # Define aa to integer table788    def aa2int(seq: str) -> List[int]:789        _aa2int = {790            "A": 1,791            "R": 2,792            "N": 3,793            "D": 4,794            "C": 5,795            "Q": 6,796            "E": 7,797            "G": 8,798            "H": 9,799            "I": 10,800            "L": 11,801            "K": 12,802            "M": 13,803            "F": 14,804            "P": 15,805            "S": 16,806            "T": 17,807            "W": 18,808            "Y": 19,809            "V": 20,810            "B": 21,811            "Z": 22,812            "X": 23,813            "*": 24,814            "-": 25,815            "?": 26,816        }817        return [_aa2int[i] for i in seq]818 819    # Create empty array to fill820    oh_array = np.zeros(shape=(26, len(input_seq)))821 822    # Load placements from aa2int823    aa_placement = aa2int(input_seq)824 825    # One-hot encode the amino acid sequence:826 827    # style nit: more pythonic to write for i in range(0, len(aa_placement)):828    for i in range(0, len(aa_placement)):829        oh_array[aa_placement[i], i] = 1830        i += 1831 832    oh_array = [oh_array]833    x = np.array(np.transpose(oh_array))834 835    y = x.astype(np.float32)836 837    y = np.reshape(y, (y.shape[0], 1, 26))838 839    # Start ICOR session using model.840    sess = rt.InferenceSession(model_path)841    input_name = sess.get_inputs()[0].name842 843    # Get prediction:844    pred_onx = sess.run(None, {input_name: y})845 846    # Get the index of the highest probability from softmax output:847    pred_indices = []848    for pred in pred_onx[0]:849        pred_indices.append(np.argmax(pred))850 851    out_str = ""852    for index in pred_indices:853        out_str += labels[index]854 855    return out_str856