OneScience-Group/CodonTransformer
07
1"""2File: CodonData.py3---------------------4Includes helper functions for preprocessing NCBI or Kazusa databases and5preparing the data for training and inference of the CodonTransformer model.6"""7 8import json9import os10import random11from typing import Dict, List, Optional, Tuple, Union12 13import pandas as pd14import python_codon_tables as pct15from Bio import SeqIO16from Bio.Seq import Seq17from sklearn.utils import shuffle as sk_shuffle18from tqdm import tqdm19 20from CodonTransformer.CodonUtils import (21 AMBIGUOUS_AMINOACID_MAP,22 AMINO2CODON_TYPE,23 AMINO_ACIDS,24 ORGANISM2ID,25 START_CODONS,26 STOP_CODONS,27 STOP_SYMBOL,28 STOP_SYMBOLS,29 ProteinConfig,30 find_pattern_in_fasta,31 get_taxonomy_id,32 sort_amino2codon_skeleton,33)34 35 36def prepare_training_data(37 dataset: Union[str, pd.DataFrame], output_file: str, shuffle: bool = True38) -> None:39 """40 Prepare a JSON dataset for training the CodonTransformer model.41 42 Input dataset should have columns below:43 - dna: str (DNA sequence)44 - protein: str (Protein sequence)45 - organism: Union[int, str] (ID or Name of the organism)46 47 The output JSON dataset will have the following format:48 {"idx": 0, "codons": "M_ATG R_AGG L_TTG L_CTA R_CGA __TAG", "organism": 51}49 {"idx": 1, "codons": "M_ATG K_AAG C_TGC F_TTT F_TTC __TAA", "organism": 59}50 51 Args:52 dataset (Union[str, pd.DataFrame]): Input dataset in CSV or DataFrame format.53 output_file (str): Path to save the output JSON dataset.54 shuffle (bool, optional): Whether to shuffle the dataset before saving.55 Defaults to True.56 57 Returns:58 None59 """60 if isinstance(dataset, str):61 dataset = pd.read_csv(dataset)62 63 required_columns = {"dna", "protein", "organism"}64 if not required_columns.issubset(dataset.columns):65 raise ValueError(f"Input dataset must have columns: {required_columns}")66 67 # Prepare the dataset for finetuning68 dataset["codons"] = dataset.apply(69 lambda row: get_merged_seq(row["protein"], row["dna"], separator="_"), axis=170 )71 72 # Replace organism str with organism id using ORGANISM2ID73 dataset["organism"] = dataset["organism"].apply(74 lambda org: process_organism(org, ORGANISM2ID)75 )76 77 # Save the dataset to a JSON file78 dataframe_to_json(dataset[["codons", "organism"]], output_file, shuffle=shuffle)79 80 81def dataframe_to_json(df: pd.DataFrame, output_file: str, shuffle: bool = True) -> None:82 """83 Convert pandas DataFrame to JSON file format suitable for training CodonTransformer.84 85 This function takes a preprocessed DataFrame and writes it to a JSON file86 where each line is a JSON object representing a single record.87 88 Args:89 df (pd.DataFrame): The input DataFrame with 'codons' and 'organism' columns.90 output_file (str): Path to the output JSON file.91 shuffle (bool, optional): Whether to shuffle the dataset before saving.92 Defaults to True.93 94 Returns:95 None96 97 Raises:98 ValueError: If the required columns are not present in the DataFrame.99 """100 required_columns = {"codons", "organism"}101 if not required_columns.issubset(df.columns):102 raise ValueError(f"DataFrame must contain columns: {required_columns}")103 104 print(f"\nStarted writing to {output_file}...")105 106 # Shuffle the DataFrame if requested107 if shuffle:108 df = sk_shuffle(df)109 110 # Write the DataFrame to a JSON file111 with open(output_file, "w") as f:112 for idx, row in tqdm(113 df.iterrows(), total=len(df), desc="Writing JSON...", unit=" records"114 ):115 doc = {"idx": idx, "codons": row["codons"], "organism": row["organism"]}116 f.write(json.dumps(doc) + "\n")117 118 print(f"\nTotal Entries Saved: {len(df)}, JSON data saved to {output_file}")119 120 121def process_organism(organism: Union[str, int], organism_to_id: Dict[str, int]) -> int:122 """123 Process and validate the organism input, converting it to a valid organism ID.124 125 This function handles both string (organism name) and integer (organism ID) inputs.126 It validates the input against a provided mapping of organism names to IDs.127 128 Args:129 organism (Union[str, int]): Input organism, either as a name (str) or ID (int).130 organism_to_id (Dict[str, int]): Dictionary mapping organism names to their131 corresponding IDs.132 133 Returns:134 int: The validated organism ID.135 136 Raises:137 ValueError: If the input is an invalid organism name or ID.138 TypeError: If the input is neither a string nor an integer.139 """140 if isinstance(organism, str):141 if organism not in organism_to_id:142 raise ValueError(f"Invalid organism name: {organism}")143 return organism_to_id[organism]144 145 elif isinstance(organism, int):146 if organism not in organism_to_id.values():147 raise ValueError(f"Invalid organism ID: {organism}")148 return organism149 150 raise TypeError(151 f"Organism must be a string or integer, not {type(organism).__name__}"152 )153 154 155def preprocess_protein_sequence(protein: str) -> str:156 """157 Preprocess a protein sequence by cleaning, standardizing, and handling158 ambiguous amino acids.159 160 Args:161 protein (str): The input protein sequence.162 163 Returns:164 str: The preprocessed protein sequence.165 166 Raises:167 ValueError: If the protein sequence is invalid or if the configuration is invalid.168 """169 if not protein:170 raise ValueError("Protein sequence is empty.")171 172 # Clean and standardize the protein sequence173 protein = (174 protein.upper().strip().replace("\n", "").replace(" ", "").replace("\t", "")175 )176 177 # Handle ambiguous amino acids based on the specified behavior178 config = ProteinConfig()179 ambiguous_aminoacid_map_override = config.get("ambiguous_aminoacid_map_override")180 ambiguous_aminoacid_behavior = config.get("ambiguous_aminoacid_behavior")181 ambiguous_aminoacid_map = AMBIGUOUS_AMINOACID_MAP.copy()182 183 for aminoacid, standard_aminoacids in ambiguous_aminoacid_map_override.items():184 ambiguous_aminoacid_map[aminoacid] = standard_aminoacids185 186 if ambiguous_aminoacid_behavior == "raise_error":187 if any(aminoacid in ambiguous_aminoacid_map for aminoacid in protein):188 raise ValueError("Ambiguous amino acids found in protein sequence.")189 elif ambiguous_aminoacid_behavior == "standardize_deterministic":190 protein = "".join(191 ambiguous_aminoacid_map.get(aminoacid, [aminoacid])[0]192 for aminoacid in protein193 )194 elif ambiguous_aminoacid_behavior == "standardize_random":195 protein = "".join(196 random.choice(ambiguous_aminoacid_map.get(aminoacid, [aminoacid]))197 for aminoacid in protein198 )199 else:200 raise ValueError(201 f"Invalid ambiguous_aminoacid_behavior: {ambiguous_aminoacid_behavior}."202 )203 204 # Check for sequence validity205 if any(aminoacid not in AMINO_ACIDS + STOP_SYMBOLS for aminoacid in protein):206 raise ValueError("Invalid characters in protein sequence.")207 208 if protein[-1] not in AMINO_ACIDS + STOP_SYMBOLS:209 raise ValueError(210 "Protein sequence must end with `*`, or `_`, or an amino acid."211 )212 213 # Replace '*' at the end of protein with STOP_SYMBOL if present214 if protein[-1] == "*":215 protein = protein[:-1] + STOP_SYMBOL216 217 # Add stop symbol to end of protein218 if protein[-1] != STOP_SYMBOL:219 protein += STOP_SYMBOL220 221 return protein222 223 224def replace_ambiguous_codons(dna: str) -> str:225 """226 Replaces ambiguous codons in a DNA sequence with "UNK".227 228 Args:229 dna (str): The DNA sequence to process.230 231 Returns:232 str: The processed DNA sequence with ambiguous codons replaced by "UNK".233 """234 result = []235 dna = dna.upper()236 237 # Check codons in DNA sequence238 for i in range(0, len(dna), 3):239 codon = dna[i : i + 3]240 241 if len(codon) == 3 and all(nucleotide in "ATCG" for nucleotide in codon):242 result.append(codon)243 else:244 result.append("UNK")245 246 return "".join(result)247 248 249def preprocess_dna_sequence(dna: str) -> str:250 """251 Cleans and preprocesses a DNA sequence by standardizing it and replacing252 ambiguous codons.253 254 Args:255 dna (str): The DNA sequence to preprocess.256 257 Returns:258 str: The cleaned and preprocessed DNA sequence.259 """260 if not dna:261 return ""262 263 # Clean and standardize the DNA sequence264 dna = dna.upper().strip().replace("\n", "").replace(" ", "").replace("\t", "")265 266 # Replace codons with ambigous nucleotides with "UNK"267 dna = replace_ambiguous_codons(dna)268 269 # Add unkown stop codon to end of DNA sequence if not present270 if dna[-3:] not in STOP_CODONS:271 dna += "UNK"272 273 return dna274 275 276def get_merged_seq(protein: str, dna: str = "", separator: str = "_") -> str:277 """278 Return the merged sequence of protein amino acids and DNA codons in the form279 of tokens separated by space, where each token is composed of an amino acid +280 separator + codon.281 282 Args:283 protein (str): Protein sequence.284 dna (str): DNA sequence.285 separator (str): Separator between amino acid and codon.286 287 Returns:288 str: Merged sequence.289 290 Example:291 >>> get_merged_seq(protein="MAV_", dna="ATGGCTGTGTAA", separator="_")292 'M_ATG A_GCT V_GTG __TAA'293 294 >>> get_merged_seq(protein="QHH_", dna="", separator="_")295 'Q_UNK H_UNK H_UNK __UNK'296 """297 merged_seq = ""298 299 # Prepare protein and dna sequences300 dna = preprocess_dna_sequence(dna)301 protein = preprocess_protein_sequence(protein)302 303 # Check if the length of protein and dna sequences are equal304 if len(dna) > 0 and len(protein) != len(dna) / 3:305 raise ValueError(306 'Length of protein (including stop symbol such as "_") and '307 "the number of codons in DNA sequence (including stop codon) "308 "must be equal."309 )310 311 # Merge protein and DNA sequences into tokens312 for i, aminoacid in enumerate(protein):313 merged_seq += f'{aminoacid}{separator}{dna[i * 3:i * 3 + 3] if dna else "UNK"} '314 315 return merged_seq.strip()316 317 318def is_correct_seq(dna: str, protein: str, stop_symbol: str = STOP_SYMBOL) -> bool:319 """320 Check if the given DNA and protein pair is correct, that is:321 1. The length of dna is divisible by 3322 2. There is an initiator codon in the beginning of dna323 3. There is only one stop codon in the sequence324 4. The only stop codon is the last codon325 326 Note since in Codon Table 3, 'TGA' is interpreted as Triptophan (W),327 there is a separate check to make sure those sequences are considered correct.328 329 Args:330 dna (str): DNA sequence.331 protein (str): Protein sequence.332 stop_symbol (str): Stop symbol.333 334 Returns:335 bool: True if the sequence is correct, False otherwise.336 """337 return (338 len(dna) % 3 == 0 # Check if DNA length is divisible by 3339 and dna[:3].upper() in START_CODONS # Check for initiator codon340 and protein[-1]341 == stop_symbol # Check if the last protein symbol is the stop symbol342 and protein.count(stop_symbol) == 1 # Check if there is only one stop symbol343 and len(set(dna))344 == 4 # Check if DNA consists of 4 unique nucleotides (A, T, C, G)345 )346 347 348def get_amino_acid_sequence(349 dna: str,350 stop_symbol: str = "_",351 codon_table: int = 1,352 return_correct_seq: bool = False,353) -> Union[str, Tuple[str, bool]]:354 """355 Return the translated protein sequence given a DNA sequence and codon table.356 357 Args:358 dna (str): DNA sequence.359 stop_symbol (str): Stop symbol.360 codon_table (int): Codon table number.361 return_correct_seq (bool): Whether to return if the sequence is correct.362 363 Returns:364 Union[str, Tuple[str, bool]]: Protein sequence and correctness flag if365 return_correct_seq is True, otherwise just the protein sequence.366 """367 dna_seq = Seq(dna).strip()368 369 # Translate the DNA sequence to a protein sequence370 protein_seq = str(371 dna_seq.translate(372 stop_symbol=stop_symbol, # Symbol to use for stop codons373 to_stop=False, # Translate the entire sequence, including any stop codons374 cds=False, # Do not assume the input is a coding sequence375 table=codon_table, # Codon table to use for translation376 )377 ).strip()378 379 return (380 protein_seq381 if not return_correct_seq382 else (protein_seq, is_correct_seq(dna_seq, protein_seq, stop_symbol))383 )384 385 386def read_fasta_file(387 input_file: str,388 save_to_file: Optional[str] = None,389 organism: str = "",390 buffer_size: int = 50000,391) -> pd.DataFrame:392 """393 Read a FASTA file of DNA sequences and convert it to a Pandas DataFrame.394 Optionally, save the DataFrame to a CSV file.395 396 Args:397 input_file (str): Path to the input FASTA file.398 save_to_file (Optional[str]): Path to save the output DataFrame. If None,399 data is only returned.400 organism (str): Name of the organism. If empty, it will be extracted from401 the FASTA description.402 buffer_size (int): Number of records to process before writing to file.403 404 Returns:405 pd.DataFrame: DataFrame containing the DNA sequences if return_dataframe406 is True, else None.407 408 Raises:409 FileNotFoundError: If the input file does not exist.410 """411 if not os.path.exists(input_file):412 raise FileNotFoundError(f"Input file not found: {input_file}")413 414 buffer = []415 columns = [416 "dna",417 "protein",418 "correct_seq",419 "organism",420 "GeneID",421 "description",422 "tokenized",423 ]424 425 # Initialize DataFrame to store all data if return_dataframe is True426 all_data = pd.DataFrame(columns=columns)427 428 with open(input_file, "r") as fasta_file:429 for record in tqdm(430 SeqIO.parse(fasta_file, "fasta"),431 desc=f"Processing {organism}",432 unit=" Records",433 ):434 dna = str(record.seq).strip().upper() # Ensure uppercase DNA sequence435 436 # Determine the organism from the record if not provided437 current_organism = organism or find_pattern_in_fasta(438 "organism", record.description439 )440 gene_id = find_pattern_in_fasta("GeneID", record.description)441 442 # Get the appropriate codon table for the organism443 codon_table = get_codon_table(current_organism)444 445 # Translate DNA to protein sequence446 protein, correct_seq = get_amino_acid_sequence(447 dna,448 stop_symbol=STOP_SYMBOL,449 codon_table=codon_table,450 return_correct_seq=True,451 )452 description = record.description.split("[", 1)[0].strip()453 tokenized = get_merged_seq(protein, dna, separator=STOP_SYMBOL)454 455 # Create a data row for the current sequence456 data_row = {457 "dna": dna,458 "protein": protein,459 "correct_seq": correct_seq,460 "organism": current_organism,461 "GeneID": gene_id,462 "description": description,463 "tokenized": tokenized,464 }465 buffer.append(data_row)466 467 # Write buffer to CSV file when buffer size is reached468 if save_to_file and len(buffer) >= buffer_size:469 write_buffer_to_csv(buffer, save_to_file, columns)470 buffer = []471 472 all_data = pd.concat(473 [all_data, pd.DataFrame([data_row])], ignore_index=True474 )475 476 # Write remaining buffer to CSV file477 if save_to_file and buffer:478 write_buffer_to_csv(buffer, save_to_file, columns)479 480 return all_data481 482 483def write_buffer_to_csv(buffer: List[Dict], output_path: str, columns: List[str]):484 """Helper function to write buffer to CSV file."""485 buffer_df = pd.DataFrame(buffer, columns=columns)486 buffer_df.to_csv(487 output_path,488 mode="a",489 header=(not os.path.exists(output_path)),490 index=True,491 )492 493 494def download_codon_frequencies_from_kazusa(495 taxonomy_id: Optional[int] = None,496 organism: Optional[str] = None,497 taxonomy_reference: Optional[str] = None,498 return_original_format: bool = False,499) -> AMINO2CODON_TYPE:500 """501 Return the codon table of the given taxonomy ID from the Kazusa Database.502 503 Args:504 taxonomy_id (Optional[int]): Taxonomy ID.505 organism (Optional[str]): Name of the organism.506 taxonomy_reference (Optional[str]): Taxonomy reference.507 return_original_format (bool): Whether to return in the original format.508 509 Returns:510 AMINO2CODON_TYPE: Codon table.511 """512 if taxonomy_reference:513 taxonomy_id = get_taxonomy_id(taxonomy_reference, organism=organism)514 515 kazusa_amino2codon = pct.get_codons_table(table_name=taxonomy_id)516 517 if return_original_format:518 return kazusa_amino2codon519 520 # Replace "*" with STOP_SYMBOL in the codon table521 kazusa_amino2codon[STOP_SYMBOL] = kazusa_amino2codon.pop("*")522 523 # Create amino2codon dictionary524 amino2codon = {525 aminoacid: (list(codon2freq.keys()), list(codon2freq.values()))526 for aminoacid, codon2freq in kazusa_amino2codon.items()527 }528 529 return sort_amino2codon_skeleton(amino2codon)530 531 532def build_amino2codon_skeleton(organism: str) -> AMINO2CODON_TYPE:533 """534 Return the empty skeleton of the amino2codon dictionary, needed for535 get_codon_frequencies.536 537 Args:538 organism (str): Name of the organism.539 540 Returns:541 AMINO2CODON_TYPE: Empty amino2codon dictionary.542 """543 amino2codon = {}544 possible_codons = [f"{i}{j}{k}" for i in "ACGT" for j in "ACGT" for k in "ACGT"]545 possible_aminoacids = get_amino_acid_sequence(546 dna="".join(possible_codons),547 codon_table=get_codon_table(organism),548 return_correct_seq=False,549 )550 551 # Initialize the amino2codon skeleton with all possible codons and set their552 # frequencies to 0553 for i, (codon, amino) in enumerate(zip(possible_codons, possible_aminoacids)):554 if amino not in amino2codon:555 amino2codon[amino] = ([], [])556 557 amino2codon[amino][0].append(codon)558 amino2codon[amino][1].append(0)559 560 # Sort the dictionary and each list of codon frequency alphabetically561 amino2codon = sort_amino2codon_skeleton(amino2codon)562 563 return amino2codon564 565 566def get_codon_frequencies(567 dna_sequences: List[str],568 protein_sequences: Optional[List[str]] = None,569 organism: Optional[str] = None,570) -> AMINO2CODON_TYPE:571 """572 Return a dictionary mapping each codon to its respective frequency based on573 the collection of DNA sequences and protein sequences.574 575 Args:576 dna_sequences (List[str]): List of DNA sequences.577 protein_sequences (Optional[List[str]]): List of protein sequences.578 organism (Optional[str]): Name of the organism.579 580 Returns:581 AMINO2CODON_TYPE: Dictionary mapping each amino acid to a tuple of codons582 and frequencies.583 """584 if organism:585 codon_table = get_codon_table(organism)586 protein_sequences = [587 get_amino_acid_sequence(588 dna, codon_table=codon_table, return_correct_seq=False589 )590 for dna in dna_sequences591 ]592 593 amino2codon = build_amino2codon_skeleton(organism)594 595 # Count the frequencies of each codon for each amino acid596 for dna, protein in zip(dna_sequences, protein_sequences):597 for i, amino in enumerate(protein):598 codon = dna[i * 3 : (i + 1) * 3]599 codon_loc = amino2codon[amino][0].index(codon)600 amino2codon[amino][1][codon_loc] += 1601 602 # Normalize codon frequencies per amino acid so they sum to 1603 amino2codon = {604 amino: (codons, [freq / (sum(frequencies) + 1e-100) for freq in frequencies])605 for amino, (codons, frequencies) in amino2codon.items()606 }607 608 return amino2codon609 610 611def get_organism_to_codon_frequencies(612 dataset: pd.DataFrame, organisms: List[str]613) -> Dict[str, AMINO2CODON_TYPE]:614 """615 Return a dictionary mapping each organism to their codon frequency distribution.616 617 Args:618 dataset (pd.DataFrame): DataFrame containing DNA sequences.619 organisms (List[str]): List of organisms.620 621 Returns:622 Dict[str, AMINO2CODON_TYPE]: Dictionary mapping each organism to its codon623 frequency distribution.624 """625 organism2frequencies = {}626 627 # Calculate codon frequencies for each organism in the dataset628 for organism in tqdm(629 organisms, desc="Calculating Codon Frequencies: ", unit="Organism"630 ):631 organism_data = dataset.loc[dataset["organism"] == organism]632 633 dna_sequences = organism_data["dna"].to_list()634 protein_sequences = organism_data["protein"].to_list()635 636 codon_frequencies = get_codon_frequencies(dna_sequences, protein_sequences)637 organism2frequencies[organism] = codon_frequencies638 639 return organism2frequencies640 641 642def get_codon_table(organism: str) -> int:643 """644 Return the appropriate NCBI codon table for a given organism.645 646 Args:647 organism (str): Name of the organism.648 649 Returns:650 int: Codon table number.651 """652 # Common codon table (Table 1) for many model organisms653 if organism in [654 "Arabidopsis thaliana",655 "Caenorhabditis elegans",656 "Chlamydomonas reinhardtii",657 "Saccharomyces cerevisiae",658 "Danio rerio",659 "Drosophila melanogaster",660 "Homo sapiens",661 "Mus musculus",662 "Nicotiana tabacum",663 "Solanum tuberosum",664 "Solanum lycopersicum",665 "Oryza sativa",666 "Glycine max",667 "Zea mays",668 ]:669 codon_table = 1670 671 # Chloroplast codon table (Table 11)672 elif organism in [673 "Chlamydomonas reinhardtii chloroplast",674 "Nicotiana tabacum chloroplast",675 ]:676 codon_table = 11677 678 # Default to Table 11 for other bacteria and archaea679 else:680 codon_table = 11681 682 return codon_table683 