OneScience-Group/CodonTransformer
07
1import random2import unittest3import warnings4 5import torch6 7from CodonTransformer.CodonData import get_amino_acid_sequence8from CodonTransformer.CodonPrediction import (9 load_model,10 load_tokenizer,11 predict_dna_sequence,12)13from CodonTransformer.CodonUtils import (14 AMINO_ACIDS,15 ORGANISM2ID,16 STOP_SYMBOLS,17 DNASequencePrediction,18)19 20 21class TestCodonPrediction(unittest.TestCase):22 @classmethod23 def setUpClass(cls):24 cls.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")25 26 # Suppress warnings about loading from HuggingFace27 for message in [28 "Tokenizer path not provided. Loading from HuggingFace.",29 "Model path not provided. Loading from HuggingFace.",30 ]:31 warnings.filterwarnings("ignore", message=message)32 33 cls.model = load_model(device=cls.device)34 cls.tokenizer = load_tokenizer()35 36 def test_predict_dna_sequence_valid_input(self):37 protein_sequence = "MWWMW"38 organism = "Escherichia coli general"39 result = predict_dna_sequence(40 protein_sequence,41 organism,42 device=self.device,43 tokenizer=self.tokenizer,44 model=self.model,45 )46 self.assertIsInstance(result.predicted_dna, str)47 self.assertTrue(48 all(nucleotide in "ATCG" for nucleotide in result.predicted_dna)49 )50 self.assertEqual(result.predicted_dna, "ATGTGGTGGATGTGGTGA")51 52 def test_predict_dna_sequence_non_deterministic(self):53 protein_sequence = "MFWY"54 organism = "Escherichia coli general"55 num_iterations = 10056 temperatures = [0.2, 0.5, 0.8]57 possible_outputs = set()58 possible_encodings_wo_stop = {59 "ATGTTTTGGTAT",60 "ATGTTCTGGTAT",61 "ATGTTTTGGTAC",62 "ATGTTCTGGTAC",63 }64 for _ in range(num_iterations):65 for temperature in temperatures:66 result = predict_dna_sequence(67 protein=protein_sequence,68 organism=organism,69 device=self.device,70 tokenizer=self.tokenizer,71 model=self.model,72 deterministic=False,73 temperature=temperature,74 )75 possible_outputs.add(result.predicted_dna[:-3]) # Remove stop codon76 77 self.assertEqual(possible_outputs, possible_encodings_wo_stop)78 79 def test_predict_dna_sequence_invalid_inputs(self):80 test_cases = [81 ("MKTZZFVLLL?", "Escherichia coli general", "invalid protein sequence"),82 ("MKTFFVLLL", "Alien $%#@!", "invalid organism code"),83 ("", "Escherichia coli general", "empty protein sequence"),84 ]85 86 for protein_sequence, organism, error_type in test_cases:87 with self.subTest(error_type=error_type):88 with self.assertRaises(ValueError):89 predict_dna_sequence(90 protein_sequence,91 organism,92 device=self.device,93 tokenizer=self.tokenizer,94 model=self.model,95 )96 97 def test_predict_dna_sequence_top_p_effect(self):98 """Test that changing top_p affects the diversity of outputs."""99 protein_sequence = "MFWY"100 organism = "Escherichia coli general"101 num_iterations = 50102 temperature = 0.5103 top_p_values = [0.8, 0.95]104 outputs_by_top_p = {top_p: set() for top_p in top_p_values}105 106 for top_p in top_p_values:107 for _ in range(num_iterations):108 result = predict_dna_sequence(109 protein=protein_sequence,110 organism=organism,111 device=self.device,112 tokenizer=self.tokenizer,113 model=self.model,114 deterministic=False,115 temperature=temperature,116 top_p=top_p,117 )118 outputs_by_top_p[top_p].add(119 result.predicted_dna[:-3]120 ) # Remove stop codon121 122 # Assert that higher top_p results in more diverse outputs123 diversity_lower_top_p = len(outputs_by_top_p[0.8])124 diversity_higher_top_p = len(outputs_by_top_p[0.95])125 self.assertGreaterEqual(126 diversity_higher_top_p,127 diversity_lower_top_p,128 "Higher top_p should result in more diverse outputs",129 )130 131 def test_predict_dna_sequence_invalid_temperature_and_top_p(self):132 """Test that invalid temperature and top_p values raise ValueError."""133 protein_sequence = "MWWMW"134 organism = "Escherichia coli general"135 invalid_params = [136 {"temperature": -0.1, "top_p": 0.95},137 {"temperature": 0, "top_p": 0.95},138 {"temperature": 0.5, "top_p": -0.1},139 {"temperature": 0.5, "top_p": 1.1},140 ]141 142 for params in invalid_params:143 with self.subTest(params=params):144 with self.assertRaises(ValueError):145 predict_dna_sequence(146 protein=protein_sequence,147 organism=organism,148 device=self.device,149 tokenizer=self.tokenizer,150 model=self.model,151 deterministic=False,152 temperature=params["temperature"],153 top_p=params["top_p"],154 )155 156 def test_predict_dna_sequence_translation_consistency(self):157 """Test that the predicted DNA translates back to the original protein."""158 protein_sequence = "MALWMRLLPLLALLALWGPDPAAAFVNQHLCGSHLVE"159 organism = "Escherichia coli general"160 result = predict_dna_sequence(161 protein=protein_sequence,162 organism=organism,163 device=self.device,164 tokenizer=self.tokenizer,165 model=self.model,166 deterministic=True,167 )168 169 # Translate predicted DNA back to protein170 translated_protein = get_amino_acid_sequence(result.predicted_dna[:-3])171 172 self.assertEqual(173 translated_protein,174 protein_sequence,175 "Translated protein does not match the original protein sequence",176 )177 178 def test_predict_dna_sequence_long_protein_sequence(self):179 """Test the function with a very long protein sequence to check performance and correctness."""180 protein_sequence = (181 "M"182 + "MALWMRLLPLLALLALWGPDPAAAFVNQHLCGSHLVEALYLVCGERGFFYTPKTRREAEDLQVGQVELGG"183 * 20184 + STOP_SYMBOLS[0]185 )186 organism = "Escherichia coli general"187 result = predict_dna_sequence(188 protein=protein_sequence,189 organism=organism,190 device=self.device,191 tokenizer=self.tokenizer,192 model=self.model,193 deterministic=True,194 )195 196 # Check that the predicted DNA translates back to the original protein197 dna_sequence = result.predicted_dna[:-3]198 translated_protein = get_amino_acid_sequence(dna_sequence)199 self.assertEqual(200 translated_protein,201 protein_sequence[:-1],202 "Translated protein does not match the original long protein sequence",203 )204 205 def test_predict_dna_sequence_edge_case_organisms(self):206 """Test the function with organism IDs at the boundaries of the mapping."""207 protein_sequence = "MWWMW"208 # Assuming ORGANISM2ID has IDs starting from 0 to N209 min_organism_id = min(ORGANISM2ID.values())210 max_organism_id = max(ORGANISM2ID.values())211 organisms = [min_organism_id, max_organism_id]212 213 for organism_id in organisms:214 with self.subTest(organism_id=organism_id):215 result = predict_dna_sequence(216 protein=protein_sequence,217 organism=organism_id,218 device=self.device,219 tokenizer=self.tokenizer,220 model=self.model,221 deterministic=True,222 )223 self.assertIsInstance(result.predicted_dna, str)224 self.assertTrue(225 all(nucleotide in "ATCG" for nucleotide in result.predicted_dna)226 )227 228 def test_predict_dna_sequence_concurrent_calls(self):229 """Test the function's behavior under concurrent execution."""230 import threading231 232 protein_sequence = "MWWMW"233 organism = "Escherichia coli general"234 results = []235 236 def call_predict():237 result = predict_dna_sequence(238 protein=protein_sequence,239 organism=organism,240 device=self.device,241 tokenizer=self.tokenizer,242 model=self.model,243 deterministic=True,244 )245 results.append(result.predicted_dna)246 247 threads = [threading.Thread(target=call_predict) for _ in range(10)]248 for thread in threads:249 thread.start()250 for thread in threads:251 thread.join()252 253 self.assertEqual(len(results), 10)254 self.assertTrue(all(dna == results[0] for dna in results))255 256 def test_predict_dna_sequence_random_seed_consistency(self):257 """Test that setting a random seed results in consistent outputs in non-deterministic mode."""258 protein_sequence = "MFWY"259 organism = "Escherichia coli general"260 temperature = 0.5261 top_p = 0.95262 torch.manual_seed(42)263 264 result1 = predict_dna_sequence(265 protein=protein_sequence,266 organism=organism,267 device=self.device,268 tokenizer=self.tokenizer,269 model=self.model,270 deterministic=False,271 temperature=temperature,272 top_p=top_p,273 )274 275 torch.manual_seed(42)276 277 result2 = predict_dna_sequence(278 protein=protein_sequence,279 organism=organism,280 device=self.device,281 tokenizer=self.tokenizer,282 model=self.model,283 deterministic=False,284 temperature=temperature,285 top_p=top_p,286 )287 288 self.assertEqual(289 result1.predicted_dna,290 result2.predicted_dna,291 "Outputs should be consistent when random seed is set",292 )293 294 def test_predict_dna_sequence_invalid_tokenizer_and_model(self):295 """Test that providing invalid tokenizer or model raises appropriate exceptions."""296 protein_sequence = "MWWMW"297 organism = "Escherichia coli general"298 299 with self.subTest("Invalid tokenizer"):300 with self.assertRaises(Exception):301 predict_dna_sequence(302 protein=protein_sequence,303 organism=organism,304 device=self.device,305 tokenizer="invalid_tokenizer_path",306 model=self.model,307 )308 309 with self.subTest("Invalid model"):310 with self.assertRaises(Exception):311 predict_dna_sequence(312 protein=protein_sequence,313 organism=organism,314 device=self.device,315 tokenizer=self.tokenizer,316 model="invalid_model_path",317 )318 319 def test_predict_dna_sequence_stop_codon_handling(self):320 """Test the function's handling of protein sequences ending with a non '_' or '*' stop symbol."""321 protein_sequence = "MWW/"322 organism = "Escherichia coli general"323 324 with self.assertRaises(ValueError):325 predict_dna_sequence(326 protein=protein_sequence,327 organism=organism,328 device=self.device,329 tokenizer=self.tokenizer,330 model=self.model,331 )332 333 def test_predict_dna_sequence_device_compatibility(self):334 """Test that the function works correctly on both CPU and GPU devices."""335 protein_sequence = "MWWMW"336 organism = "Escherichia coli general"337 338 devices = [torch.device("cpu")]339 if torch.cuda.is_available():340 devices.append(torch.device("cuda"))341 342 for device in devices:343 with self.subTest(device=device):344 result = predict_dna_sequence(345 protein=protein_sequence,346 organism=organism,347 device=device,348 tokenizer=self.tokenizer,349 model=self.model,350 deterministic=True,351 )352 self.assertIsInstance(result.predicted_dna, str)353 self.assertTrue(354 all(nucleotide in "ATCG" for nucleotide in result.predicted_dna)355 )356 357 def test_predict_dna_sequence_random_proteins(self):358 """Test random proteins to ensure translated DNA matches the original protein."""359 organism = "Escherichia coli general"360 num_tests = 200361 362 for _ in range(num_tests):363 # Generate a random protein sequence of random length between 10 and 50364 protein_length = random.randint(10, 500)365 protein_sequence = "M" + "".join(366 random.choices(AMINO_ACIDS, k=protein_length - 1)367 )368 protein_sequence += random.choice(STOP_SYMBOLS)369 370 result = predict_dna_sequence(371 protein=protein_sequence,372 organism=organism,373 device=self.device,374 tokenizer=self.tokenizer,375 model=self.model,376 deterministic=True,377 )378 379 # Remove stop codon from predicted DNA380 dna_sequence = result.predicted_dna[:-3]381 382 # Translate predicted DNA back to protein383 translated_protein = get_amino_acid_sequence(dna_sequence)384 self.assertEqual(385 translated_protein,386 protein_sequence[:-1], # Remove stop symbol387 f"Translated protein does not match the original protein sequence for protein: {protein_sequence}",388 )389 390 def test_predict_dna_sequence_long_protein_over_max_length(self):391 """Test that the model handles protein sequences longer than 2048 amino acids."""392 # Create a protein sequence longer than 2048 amino acids393 base_sequence = (394 "MALWMRLLPLLALLALWGPDPAAAFVNQHLCGSHLVEALYLVCGERGFFYTPKTRREAEDLQVGQVELGG"395 )396 protein_sequence = base_sequence * 100 # Length > 2048 amino acids397 organism = "Escherichia coli general"398 399 result = predict_dna_sequence(400 protein=protein_sequence,401 organism=organism,402 device=self.device,403 tokenizer=self.tokenizer,404 model=self.model,405 deterministic=True,406 )407 408 # Remove stop codon from predicted DNA409 dna_sequence = result.predicted_dna[:-3]410 translated_protein = get_amino_acid_sequence(dna_sequence)411 412 # Due to potential model limitations, compare up to the model's max supported length413 max_length = len(translated_protein)414 self.assertEqual(415 translated_protein[:max_length],416 protein_sequence[:max_length],417 "Translated protein does not match the original protein sequence up to the maximum length supported.",418 )419 420 def test_predict_dna_sequence_multi_output(self):421 """Test that the function returns multiple sequences when num_sequences > 1."""422 protein_sequence = "MFQLLAPWY"423 organism = "Escherichia coli general"424 num_sequences = 20425 426 result = predict_dna_sequence(427 protein=protein_sequence,428 organism=organism,429 device=self.device,430 tokenizer=self.tokenizer,431 model=self.model,432 deterministic=False,433 num_sequences=num_sequences,434 )435 436 self.assertIsInstance(result, list)437 self.assertEqual(len(result), num_sequences)438 439 for prediction in result:440 self.assertIsInstance(prediction, DNASequencePrediction)441 self.assertTrue(442 all(nucleotide in "ATCG" for nucleotide in prediction.predicted_dna)443 )444 445 # Check that all predicted DNA sequences translate back to the original protein446 translated_protein = get_amino_acid_sequence(prediction.predicted_dna[:-3])447 self.assertEqual(translated_protein, protein_sequence)448 449 def test_predict_dna_sequence_deterministic_multi_raises_error(self):450 """Test that requesting multiple sequences in deterministic mode raises an error."""451 protein_sequence = "MFWY"452 organism = "Escherichia coli general"453 454 with self.assertRaises(ValueError):455 predict_dna_sequence(456 protein=protein_sequence,457 organism=organism,458 device=self.device,459 tokenizer=self.tokenizer,460 model=self.model,461 deterministic=True,462 num_sequences=3,463 )464 465 def test_predict_dna_sequence_multi_diversity(self):466 """Test that multiple sequences generated are diverse."""467 protein_sequence = "MFWYMFWY"468 organism = "Escherichia coli general"469 num_sequences = 10470 471 result = predict_dna_sequence(472 protein=protein_sequence,473 organism=organism,474 device=self.device,475 tokenizer=self.tokenizer,476 model=self.model,477 deterministic=False,478 num_sequences=num_sequences,479 temperature=0.8,480 )481 482 unique_sequences = set(prediction.predicted_dna for prediction in result)483 484 self.assertGreater(485 len(unique_sequences),486 2,487 "Multiple sequence generation should produce diverse results",488 )489 490 # Check that all sequences are valid translations of the input protein491 for prediction in result:492 translated_protein = get_amino_acid_sequence(prediction.predicted_dna[:-3])493 self.assertEqual(translated_protein, protein_sequence)494 495 def test_predict_dna_sequence_match_protein_repetitive(self):496 """Test that match_protein=True correctly handles highly repetitive and unconventional sequences."""497 test_sequences = (498 "QQQQQQQQQQQQQQQQ_",499 "KRKRKRKRKRKRKRKR_",500 "PGPGPGPGPGPGPGPG_",501 "DEDEDEDEDEDEDEDEDE_",502 "M_M_M_M_M_",503 "MMMMMMMMMM_",504 "WWWWWWWWWW_",505 "CCCCCCCCCC_",506 "MWCHMWCHMWCH_",507 "Q_QQ_QQQ_QQQQ_",508 "MWMWMWMWMWMW_",509 "CCCHHHMMMWWW_",510 "_",511 "M_",512 "MGWC_",513 )514 515 organism = "Homo sapiens"516 517 for protein_sequence in test_sequences:518 # Generate sequence with match_protein=True519 result = predict_dna_sequence(520 protein=protein_sequence,521 organism=organism,522 device=self.device,523 tokenizer=self.tokenizer,524 model=self.model,525 deterministic=False,526 temperature=20, # High temperature to test protein matching527 match_protein=True,528 )529 530 dna_sequence = result.predicted_dna531 translated_protein = get_amino_acid_sequence(dna_sequence)532 533 self.assertEqual(534 translated_protein,535 protein_sequence,536 f"Translated protein must match original when match_protein=True. Failed for sequence: {protein_sequence}",537 )538 539 def test_predict_dna_sequence_match_protein_rare_amino_acids(self):540 """Test match_protein with rare amino acids that have limited codon options."""541 # Methionine (M) and Tryptophan (W) have only one codon each542 # While Leucine (L) has 6 codons - testing contrast543 protein_sequence = "MWLLLMWLLL"544 organism = "Escherichia coli general"545 546 # Run multiple predictions547 results = []548 num_iterations = 10549 550 for _ in range(num_iterations):551 result = predict_dna_sequence(552 protein=protein_sequence,553 organism=organism,554 device=self.device,555 tokenizer=self.tokenizer,556 model=self.model,557 deterministic=False,558 temperature=20, # High temperature to test protein matching559 match_protein=True,560 )561 results.append(result.predicted_dna)562 563 # Check all sequences564 for dna_sequence in results:565 # Verify M always uses ATG566 m_positions = [0, 5] # Known positions of M in sequence567 for pos in m_positions:568 self.assertEqual(569 dna_sequence[pos * 3 : (pos + 1) * 3],570 "ATG",571 "Methionine must use ATG codon.",572 )573 574 # Verify W always uses TGG575 w_positions = [1, 6] # Known positions of W in sequence576 for pos in w_positions:577 self.assertEqual(578 dna_sequence[pos * 3 : (pos + 1) * 3],579 "TGG",580 "Tryptophan must use TGG codon.",581 )582 583 # Verify all L codons are valid584 l_positions = [2, 3, 4, 7, 8, 9] # Known positions of L in sequence585 l_codons = [dna_sequence[pos * 3 : (pos + 1) * 3] for pos in l_positions]586 valid_l_codons = {"TTA", "TTG", "CTT", "CTC", "CTA", "CTG"}587 self.assertTrue(588 all(codon in valid_l_codons for codon in l_codons),589 "All Leucine codons must be valid",590 )591 592 593if __name__ == "__main__":594 unittest.main()595 