OneScience-Group/CodonTransformer
07
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 