Download classifiers/scoring_functions.py from atombio/mrnautilus: direct link, hf CLI and curl.
- Browser
- Download file 10.6 kB
-
https://huggingface.co/atombio/mrnautilus/resolve/main/classifiers/scoring_functions.py
- Command line
-
hf download hf://atombio/mrnautilus/classifiers/scoring_functions.py
-
curl -L -o scoring_functions.py https://huggingface.co/atombio/mrnautilus/resolve/main/classifiers/scoring_functions.py
10.6 kB
| import numpy as np | |
| import torch | |
| from classifiers.hl.hl import HalfLife | |
| from classifiers.pa.translation_rate import TranslationRate | |
| from classifiers.te.riboseq import RiboSeq | |
| from src.mrna.tasks.lm.mdlm import Diffusion | |
| class ScoringFunctions: | |
| def __init__(self, device, five_utr_ind, cds_length, model=None, ckpt_path=None, conf_path=None, score_func_names=None): | |
| """ | |
| Class for generating score vectors given generated sequence | |
| Args: | |
| score_func_names: list of scoring function names to be evaluated | |
| score_weights: weights to scale scores (default: 1) | |
| target_protein: sequence of target protein binder | |
| """ | |
| if score_func_names is None: | |
| self.score_func_names = [] | |
| else: | |
| self.score_func_names = score_func_names | |
| self.device = device | |
| halflife = HalfLife() | |
| translation_rate = TranslationRate() | |
| ribo_seq = RiboSeq() | |
| self.five_utr_ind = five_utr_ind | |
| self.three_utr_ind = five_utr_ind + cds_length | |
| self.all_funcs = {'halflife': halflife, | |
| 'translation_rate': translation_rate, | |
| 'ribo_seq': ribo_seq | |
| } | |
| self.all_funcs = {k: v for k, v in self.all_funcs.items() if k in self.score_func_names} | |
| if model is not None: | |
| self.model = model | |
| else: | |
| self.model = Diffusion.load_from_pretrained(ckpt_path, conf_path) | |
| self.model.eval().to(self.device) # Set to evaluation mode | |
| def generate_embeddings(self, seqs): | |
| """ | |
| Converts a list of RNA sequences into hidden state embeddings using a pretrained DiffusionRNALanguageModel. | |
| Args: | |
| seqs (List[str]): List of RNA sequences (e.g., ['ACGU...', 'UGCA...']). | |
| model_ckpt_path (str): Path to the model checkpoint. | |
| Returns: | |
| np.ndarray: Array of shape (N, L, d_model), where | |
| N = number of sequences, | |
| L = sequence length, | |
| d_model = embedding dimension. | |
| """ | |
| # Load model | |
| five_utrs = [seq[:self.five_utr_ind] for seq in seqs] | |
| cds = [seq[self.five_utr_ind:self.three_utr_ind] for seq in seqs] | |
| three_utrs = [seq[self.three_utr_ind:] for seq in seqs] | |
| # Tokenize sequences | |
| tokens = torch.tensor(self.model.tokenizer.batch_tokenize(five_utrs, cds, three_utrs), dtype=torch.int64, device=self.device) | |
| with torch.no_grad(): | |
| with torch.autocast(device_type='cuda', dtype=torch.float16): | |
| hidden = self.model.net(tokens)['last_hidden_state'] | |
| return hidden.mean(dim=1).cpu().numpy() | |
| def forward(self, input_seqs): | |
| scores = [] | |
| input_embs = self.generate_embeddings(input_seqs) | |
| for i, score_func in enumerate(self.score_func_names): | |
| if score_func == 'gc_content': | |
| score = np.array([get_gc_content(seq) for seq in input_seqs]) | |
| elif score_func == 'sequence_complexity': | |
| score = np.array([idt_sequence_complexity(seq) for seq in input_seqs]) | |
| else: | |
| score = self.all_funcs[score_func](input_embs = input_embs) | |
| scores.append(score) | |
| # convert to numpy arrays with shape (num_sequences, num_functions) | |
| scores = np.float32(scores).T | |
| return scores | |
| def __call__(self, input_seqs: list): | |
| return self.forward(input_seqs) | |
| def idt_sequence_complexity(dna: str, window: int = 100, step: int = 50, homopolymer_threshold: int = 6) -> float: | |
| """ | |
| Calculate a simple sequence complexity score inspired by IDT's approach. | |
| It penalizes low complexity regions with homopolymers and repeats. | |
| Args: | |
| dna (str): DNA sequence (A,T,C,G). | |
| window (int): Sliding window size. | |
| step (int): Step size for sliding. | |
| homopolymer_threshold (int): Minimum homopolymer length to count as low complexity. | |
| Returns: | |
| float: Complexity score (0 to 1), higher is more complex. | |
| """ | |
| dna = dna.upper() | |
| def homopolymer_fraction(seq: str) -> float: | |
| """Fraction of sequence covered by homopolymers longer than threshold.""" | |
| count = 0 | |
| current_char = '' | |
| current_run = 0 | |
| for base in seq: | |
| if base == current_char: | |
| current_run += 1 | |
| else: | |
| if current_run >= homopolymer_threshold: | |
| count += current_run | |
| current_char = base | |
| current_run = 1 | |
| # Check last run | |
| if current_run >= homopolymer_threshold: | |
| count += current_run | |
| return count / len(seq) | |
| def simple_repeat_fraction(seq: str, k: int = 4) -> float: | |
| """Fraction of seq that consists of repeated k-mers.""" | |
| if len(seq) < k: | |
| return 0.0 | |
| kmers = [seq[i:i+k] for i in range(len(seq) - k + 1)] | |
| unique_kmers = set(kmers) | |
| repeat_count = sum(kmers.count(kmer) - 1 for kmer in unique_kmers if kmers.count(kmer) > 1) | |
| return repeat_count / len(seq) | |
| scores = [] | |
| for i in range(0, len(dna) - window + 1, step): | |
| sub_seq = dna[i:i+window] | |
| hp_frac = homopolymer_fraction(sub_seq) | |
| rep_frac = simple_repeat_fraction(sub_seq) | |
| # Combine penalties (you can tune weights if desired) | |
| penalty = (hp_frac + rep_frac) / 2 | |
| # Complexity = 1 - penalty | |
| complexity = max(0.0, 1.0 - penalty) | |
| scores.append(complexity) | |
| if not scores: | |
| # Fallback for short sequences | |
| hp_frac = homopolymer_fraction(dna) | |
| rep_frac = simple_repeat_fraction(dna) | |
| penalty = (hp_frac + rep_frac) / 2 | |
| return max(0.0, 1.0 - penalty) | |
| return sum(scores) / len(scores) | |
| def get_gc_content(dna:str) -> float: | |
| """ | |
| Calculate the GC content of a DNA sequence. | |
| Args: | |
| dna (str): The DNA sequence. | |
| Returns: | |
| float: The GC content as a percentage. | |
| """ | |
| if len(dna) == 0: | |
| return 0.0 | |
| gc_count = sum(1 for base in dna if base in 'GCgc') | |
| gc_content = (gc_count / len(dna)) * 100 | |
| return gc_content <= 0.62 | |
| # Write a function that finds the length of the longest consecutive sequence of the same character in a string | |
| def longest_consecutive_sequence(dna: str) -> int: | |
| """ | |
| Find the length of the longest consecutive sequence of the same character in a DNA sequence. | |
| Args: | |
| s (str): The input DNA sequence. | |
| Returns: | |
| int: The length of the longest consecutive sequence. | |
| """ | |
| max_length = 0 | |
| current_length = 1 | |
| for i in range(1, len(dna)): | |
| if dna[i] == dna[i - 1]: | |
| current_length += 1 | |
| else: | |
| max_length = max(max_length, current_length) | |
| current_length = 1 | |
| max_length = max(max_length, current_length) | |
| return 1 if max_length < 6 else 0 | |
| def unittest(): | |
| device = 'cuda' | |
| scoring = ScoringFunctions(device, score_func_names=['halflife', | |
| 'translation_rate', | |
| 'ribosome']) | |
| seq = ['ATTAAAGGTTTATACCTTCCCAGGTAACAAACCAACCAACTTTCGATCTCTTGTAGATCTGTTCTCTAAACGAACTTTAAAATCTGTGTGGCTGTCACTCGGCTGCATGCTTAGTGCACTCACGCAGTATAATTAATAACTAATTACTGTCGTTGACAGGACACGAGTAACTCGTCTATCTTCTGCAGGCTGCTTACGGTTTCGTCCGTGTTGCAGCCGATCATCAGCACATCTAGGTTTCGTCCGGGTGTGACCGAAAGGTAAGATGAGTAAAGGAGAAGAACTTTTCACTGGAGTTGTCCCAATTCTTGTTGAATTAGATGGCGATGTTAATGGGCAAAAATTCTCTGTCAGTGGAGAGGGTGAAGGTGATGCAACATACGGAAAACTTACCCTTAAATTTATTTGCACTACTGGGAAGCTACCTGTTCCATGGCCAACACTTGTCACTACTTTCTCTTATGGTGTTCAATGCTTTTCAAGATACCCAGATCATATGAAACAGCATGACTTTTTCAAGAGTGCCATGCCCGAAGGTTATGTACAGGAAAGAACTATATTTTACAAAGATGACGGGAACTACAAGACACGTGCTGAAGTCAAGTTTGAAGGTGATACCCTTGTTAATAGAATCGAGTTAAAAGGTATTGATTTTAAAGAAGATGGAAACATTCTTGGACACAAAATGGAATACAACTATAACTCACATAATGTATACATCATGGCAGACAAACCAAAGAATGGAATCAAAGTTAACTTCAAAATTAGACACAACATTAAAGATGGAAGCGTTCAATTAGCAGACCATTATCAACAAAATACTCCAATTGGCGATGGCCCTGTCCTTTTACCAGACAACCATTACCTGTCCACACAATCTGCCCTTTCCAAAGATCCCAACGAAAAGAGAGATCACATGATCCTTCTTGAGTTTGTAACAGCTGCTGGGATTACACATGGCATGGATGAACTATACAAATAACAATCTTTAATCAGTGTGTAACATTAGGGAGGACTTGAAAGAGCCACCACATTTTCACCGAGGCCACGCGGAGTACGATCGAGTGTACAGTGAACAATGCTAGGGAGAGCTGCCTATATGGAAGAGCCCTAATGTGTAAAATTAATTTTAGTAGTGCTATCCCCATGTGATTTTAATAGCTTCTTAGGAGAATGACAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA', | |
| 'TACACACGAATAAAAGATAACAAAGATGAGTAAAGGAGAAGAACTTTTCACTGGAGTTGTCCCAATTCTTGTTGAATTAGATGGCGATGTTAATGGGCAAAAATTCTCTGTCAGTGGAGAGGGTGAAGGTGATGCAACATACGGAAAACTTACCCTTAAATTTATTTGCACTACTGGGAAGCTACCTGTTCCATGGCCAACACTTGTCACTACTTTCTCTTATGGTGTTCAATGCTTTTCAAGATACCCAGATCATATGAAACAGCATGACTTTTTCAAGAGTGCCATGCCCGAAGGTTATGTACAGGAAAGAACTATATTTTACAAAGATGACGGGAACTACAAGACACGTGCTGAAGTCAAGTTTGAAGGTGATACCCTTGTTAATAGAATCGAGTTAAAAGGTATTGATTTTAAAGAAGATGGAAACATTCTTGGACACAAAATGGAATACAACTATAACTCACATAATGTATACATCATGGCAGACAAACCAAAGAATGGAATCAAAGTTAACTTCAAAATTAGACACAACATTAAAGATGGAAGCGTTCAATTAGCAGACCATTATCAACAAAATACTCCAATTGGCGATGGCCCTGTCCTTTTACCAGACAACCATTACCTGTCCACACAATCTGCCCTTTCCAAAGATCCCAACGAAAAGAGAGATCACATGATCCTTCTTGAGTTTGTAACAGCTGCTGGGATTACACATGGCATGGATGAACTATACAAATAAATGTCCAGACTTCCAATTGACACTAAAGTGTCCGAACAATTACTAAATTCTCAGGGTTCCTGGTTAAATTCAGGCTGAGACTTTATTTATATATTTATAGATTCATTAAAATTTTATGAATAATTTATTGATGTTATTAATAGGGGCTATTTTCTTATTAAATAGGCTACTGGAGTGTAT', | |
| 'CGCCTGCCTGAATCTGTTCTGCCCCCTCCCCACCCATTTCACCACCACCATGACACCGGGCACCCAGTCTCCTTTCTTCCTGCTGCTGCTCCTCACAGTGCTTACAGCTACCACAGCCCCTAAACCCGCAACAGTTGTTACGGGTTCTGGTCATGCAAGCTCTACCCCAGGTGGAGAAAAGGAGACTTCGGCTACCCAGAGAAGTTCAGTGCCCAGCTCTACTGAGAAGAATGCTTTTAATTCCTCTCTGGAAGATCCCAGCACCGACTACTACCAAGAGCTGCAGAGAGACATTTCTGAAATGTTTTTGCAGATTTATAAACAAGGGGGTTTTCTGGGCCTCTCCAATATTAAGTTCAGGCCAGGATCTGTGGTGGTACAATTGACTCTGGCCTTCCGAGAAGGTACCATCAATGTCCACGACGTGGAGACACAGTTCAATCAGTATAAAACGGAAGCAGCCTCTCGATATAACCTGACGATCTCAGACGTCAGCGTGAGTGATGTGCCATTTCCTTTCTCTGCCCAGTCTGGGGCTGGGGTGCCAGGCTGGGGCATCGCGCTGCTGGTGCTGGTCTGTGTTCTGGTTGCGCTGGCCATTGTCTATCTCATTGCCTTGGCTGTCTGTCAGTGCCGCCGAAAGAACTACGGGCAGCTGGACATCTTTCCAGCCCGGGATACCTACCATCCTATGAGCGAGTACCCCACCTACCACACCCATGGGCGCTATGTGCCCCCTAGCAGTACCGATCGTAGCCCCTATGAGAAGGTTTCTGCAGGTAATGGTGGCAGCAGCCTCTCTTACACAAACCCAGCAGTGGCAGCCACTTCTGCCAACTTGTAGGGGCACGTCGCCCGCTGAGCTGAGTGGCCAGCCAGTGCCATTCCACTCCACTCAGGTTCTTCAGGGCCAGAGCCCCTGCACCCTGTTTGGGCTGGTGAGCTGGGAGTTCAGGTGGGCTGCTCACAGCCTCCTTCAGAGGCCCCACCAATTTCTCGGACA'] | |
| scores = scoring(input_seqs=seq) | |
| print(scores) | |
| print(len(scores)) | |
| if __name__ == '__main__': | |
| unittest() |