Team Ai
Modelpublic

OneScience-Group/CodonTransformer

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes7downloads
CodonData.py683 linesDownload Raw Back to CodonTransformer
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