Team Ai
Modelpublic

Taykhoom/CodonBERT

sourceHugging Faceotherupdated 2mo agoView on Hugging Face
0likes95downloads
tokenization_codonbert.py89 linesDownload Raw Back to root
1import warnings2import numpy as np3from transformers import BertTokenizer4 5 6class CodonBertTokenizer(BertTokenizer):7    """BertTokenizer that auto-converts nucleotide sequences to codon-level tokens.8 9    Raw nucleotide input is normalized (T->U, uppercase, whitespace stripped),10    then split into non-overlapping 3-mer codons before vocab lookup. Trailing11    1-2 nucleotides that do not form a complete codon are dropped.12 13    eos_token is aliased to sep_token ("[SEP]") so that pooling code that14    excludes both CLS and EOS/SEP positions works correctly.15 16    Standard usage (raw nucleotides):17        tokenizer("AUGAAAGGG")18        tokenizer(["AUGAAAGGG", "AUGUUUCCC"], return_tensors="pt", padding=True)19 20    CDS-aware usage (full mRNA + CDS track -> extract CDS, chunk, encode):21        tokenizer.batch_encode_with_cds(22            ["NNNATGAAAGGGNN"],23            cds=[np.array([0,0,0,1,0,0,1,0,0,1,0,0,0,0])],24            return_tensors="pt",25            padding=True,26        )27 28    """29 30    def __init__(self, *args, **kwargs):31        kwargs.setdefault("eos_token", "[SEP]")32        super().__init__(*args, **kwargs)33 34    def _tokenize(self, text, split_special_tokens=False):35        seq = "".join(text.split()).upper().replace("T", "U")36        n = len(seq) - len(seq) % 337        return [seq[i:i + 3] for i in range(0, n, 3)]38 39    @staticmethod40    def _extract_cds(sequence, cds):41        if sum(cds) == 0:42            warnings.warn("No CDS found. Returning truncated sequence.")43            n = len(sequence) - len(sequence) % 344            return sequence[:n]45        first = int(np.argmax(cds == 1))46        last = int(len(cds) - 1 - np.argmax(np.flip(cds) == 1)) + 247        proposed = sequence[first:last + 1]48        if len(proposed) % 3 != 0:49            warnings.warn("Irregular CDS. Returning truncated sequence.")50            return proposed[:-(len(proposed) % 3)]51        return proposed52 53    def batch_encode_with_cds(self, sequences, cds_tracks, max_length=None, **kwargs):54        """Encode a batch of raw mRNA sequences using CDS-aware preprocessing.55 56        Args:57            sequences: List of raw nucleotide strings.58            cds_tracks: List of numpy arrays (one per sequence). Non-zero values59                mark the first nucleotide of each codon in the CDS region.60            max_length: Max content codon-tokens per chunk (special tokens NOT61                counted). Defaults to model_max_length - 2.62            **kwargs: Forwarded to batch_encode_plus (e.g. return_tensors, padding).63 64        Returns:65            (BatchEncoding, chunk_counts): chunk_counts[i] is the number of66            chunks produced from sequence i.67        """68        budget_codons = max_length or (self.model_max_length - 2)69        budget_nt = budget_codons * 370 71        all_strings = []72        chunk_counts = []73 74        for seq, cds in zip(sequences, cds_tracks):75            seq = seq.replace("T", "U").replace("t", "u").upper()76            cds_seq = self._extract_cds(seq, np.asarray(cds))77            n = len(cds_seq)78            chunks = []79            for i in range(0, max(n, 1), budget_nt):80                chunk = cds_seq[i:i + budget_nt]81                chunk = chunk[:len(chunk) - len(chunk) % 3]82                if chunk:83                    chunks.append(chunk)84            all_strings.extend(chunks or [""])85            chunk_counts.append(len(chunks) or 1)86 87        enc = self.batch_encode_plus(all_strings, **kwargs)88        return enc, chunk_counts89