mrnautilus / classifiers /scoring_functions.py
sawanp813's picture
Initial commit
e781e66
Raw History Blame Contribute Delete
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()