Team Ai
Apppublic

ujjwalpardeshi/pytorch-training-debugger

sourceHugging Faceupdated 6mo agoView on Hugging Face
2likes
test_simulation_extended.py99 linesDownload Raw Back to tests
1"""Extended simulation tests — adapted for real mini-training curves."""2 3from __future__ import annotations4 5from ml_training_debugger.scenarios import sample_scenario6from ml_training_debugger.simulation import (7    gen_data_batch_stats,8    gen_loss_history,9    gen_val_accuracy_history,10    gen_val_loss_history,11)12 13 14class TestVanishingGradients:15    def test_loss_barely_decreases(self):16        s = sample_scenario("task_002", seed=42)17        hist = gen_loss_history(s)18        assert len(hist) == 2019 20    def test_val_acc_low(self):21        s = sample_scenario("task_002", seed=42)22        hist = gen_val_accuracy_history(s)23        assert len(hist) == 2024 25    def test_val_loss_present(self):26        s = sample_scenario("task_002", seed=42)27        hist = gen_val_loss_history(s)28        assert len(hist) == 2029 30 31class TestOverfitting:32    def test_loss_history_present(self):33        s = sample_scenario("task_004", seed=42)34        hist = gen_loss_history(s)35        assert len(hist) == 2036 37    def test_val_acc_present(self):38        s = sample_scenario("task_004", seed=42)39        hist = gen_val_accuracy_history(s)40        assert len(hist) == 2041 42    def test_val_loss_present(self):43        s = sample_scenario("task_004", seed=42)44        hist = gen_val_loss_history(s)45        assert len(hist) == 2046 47    def test_data_batch_stats_clean(self):48        s = sample_scenario("task_004", seed=42)49        stats = gen_data_batch_stats(s)50        assert stats["class_overlap_score"] == 0.051        assert stats["duplicate_ratio"] == 0.052 53 54class TestCodeBug:55    def test_loss_history(self):56        s = sample_scenario("task_006", seed=42)57        hist = gen_loss_history(s)58        assert len(hist) == 2059 60    def test_val_acc(self):61        s = sample_scenario("task_006", seed=42)62        hist = gen_val_accuracy_history(s)63        assert len(hist) == 2064 65    def test_val_loss(self):66        s = sample_scenario("task_006", seed=42)67        hist = gen_val_loss_history(s)68        assert len(hist) == 2069 70 71class TestBatchNormEval:72    def test_val_loss_present(self):73        s = sample_scenario("task_005", seed=42)74        hist = gen_val_loss_history(s)75        assert len(hist) == 2076 77    def test_val_acc_near_zero(self):78        s = sample_scenario("task_005", seed=42)79        hist = gen_val_accuracy_history(s)80        # BatchNorm eval mode makes learning very poor81        assert len(hist) == 2082 83 84class TestSchedulerMisconfigured:85    def test_loss_history(self):86        s = sample_scenario("task_007", seed=42)87        hist = gen_loss_history(s)88        assert len(hist) == 2089 90    def test_val_acc(self):91        s = sample_scenario("task_007", seed=42)92        hist = gen_val_accuracy_history(s)93        assert len(hist) == 2094 95    def test_val_loss(self):96        s = sample_scenario("task_007", seed=42)97        hist = gen_val_loss_history(s)98        assert len(hist) == 2099