| 83 | |
| 84 | @require_torch |
| 85 | class ScheduleInitTest(unittest.TestCase): |
| 86 | m = torch.nn.Linear(50, 50) if is_torch_available() else None |
| 87 | optimizer = AdamW(m.parameters(), lr=10.0) if is_torch_available() else None |
| 88 | num_steps = 10 |
| 89 | |
| 90 | def assertListAlmostEqual(self, list1, list2, tol): |
| 91 | self.assertEqual(len(list1), len(list2)) |
| 92 | for a, b in zip(list1, list2): |
| 93 | self.assertAlmostEqual(a, b, delta=tol) |
| 94 | |
| 95 | def test_constant_scheduler(self): |
| 96 | scheduler = get_constant_schedule(self.optimizer) |
| 97 | lrs = unwrap_schedule(scheduler, self.num_steps) |
| 98 | expected_learning_rates = [10.0] * self.num_steps |
| 99 | self.assertEqual(len(lrs[0]), 1) |
| 100 | self.assertListEqual([l[0] for l in lrs], expected_learning_rates) |
| 101 | |
| 102 | scheduler = get_constant_schedule(self.optimizer) |
| 103 | lrs_2 = unwrap_and_save_reload_schedule(scheduler, self.num_steps) |
| 104 | self.assertListEqual([l[0] for l in lrs], [l[0] for l in lrs_2]) |
| 105 | |
| 106 | def test_warmup_constant_scheduler(self): |
| 107 | scheduler = get_constant_schedule_with_warmup(self.optimizer, num_warmup_steps=4) |
| 108 | lrs = unwrap_schedule(scheduler, self.num_steps) |
| 109 | expected_learning_rates = [2.5, 5.0, 7.5, 10.0, 10.0, 10.0, 10.0, 10.0, 10.0, 10.0] |
| 110 | self.assertEqual(len(lrs[0]), 1) |
| 111 | self.assertListEqual([l[0] for l in lrs], expected_learning_rates) |
| 112 | |
| 113 | scheduler = get_constant_schedule_with_warmup(self.optimizer, num_warmup_steps=4) |
| 114 | lrs_2 = unwrap_and_save_reload_schedule(scheduler, self.num_steps) |
| 115 | self.assertListEqual([l[0] for l in lrs], [l[0] for l in lrs_2]) |
| 116 | |
| 117 | def test_warmup_linear_scheduler(self): |
| 118 | scheduler = get_linear_schedule_with_warmup(self.optimizer, num_warmup_steps=2, num_training_steps=10) |
| 119 | lrs = unwrap_schedule(scheduler, self.num_steps) |
| 120 | expected_learning_rates = [5.0, 10.0, 8.75, 7.5, 6.25, 5.0, 3.75, 2.5, 1.25, 0.0] |
| 121 | self.assertEqual(len(lrs[0]), 1) |
| 122 | self.assertListEqual([l[0] for l in lrs], expected_learning_rates) |
| 123 | |
| 124 | scheduler = get_linear_schedule_with_warmup(self.optimizer, num_warmup_steps=2, num_training_steps=10) |
| 125 | lrs_2 = unwrap_and_save_reload_schedule(scheduler, self.num_steps) |
| 126 | self.assertListEqual([l[0] for l in lrs], [l[0] for l in lrs_2]) |
| 127 | |
| 128 | def test_warmup_cosine_scheduler(self): |
| 129 | scheduler = get_cosine_schedule_with_warmup(self.optimizer, num_warmup_steps=2, num_training_steps=10) |
| 130 | lrs = unwrap_schedule(scheduler, self.num_steps) |
| 131 | expected_learning_rates = [5.0, 10.0, 9.61, 8.53, 6.91, 5.0, 3.08, 1.46, 0.38, 0.0] |
| 132 | self.assertEqual(len(lrs[0]), 1) |
| 133 | self.assertListAlmostEqual([l[0] for l in lrs], expected_learning_rates, tol=1e-2) |
| 134 | |
| 135 | scheduler = get_cosine_schedule_with_warmup(self.optimizer, num_warmup_steps=2, num_training_steps=10) |
| 136 | lrs_2 = unwrap_and_save_reload_schedule(scheduler, self.num_steps) |
| 137 | self.assertListEqual([l[0] for l in lrs], [l[0] for l in lrs_2]) |
| 138 | |
| 139 | def test_warmup_cosine_hard_restart_scheduler(self): |
| 140 | scheduler = get_cosine_with_hard_restarts_schedule_with_warmup( |
| 141 | self.optimizer, num_warmup_steps=2, num_cycles=2, num_training_steps=10 |
| 142 | ) |
nothing calls this directly
no test coverage detected