Taykhoom/CodonBERT
095
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 