Team Ai
Modelpublic

OneScience-Group/CodonTransformer

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes7downloads
test_CodonPrediction.py595 linesDownload Raw Back to tests
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