ujjwalpardeshi/pytorch-training-debugger
2
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 