MCPcopy Create free account
hub / github.com/huggingface/transformers / ScheduleInitTest

Class ScheduleInitTest

tests/test_optimization.py:85–152  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

83
84@require_torch
85class 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 )

Callers

nothing calls this directly

Calls 2

is_torch_availableFunction · 0.90
AdamWClass · 0.90

Tested by

no test coverage detected