| 193 | self.assertEqual(m.state["hello"], "world") |
| 194 | |
| 195 | def test_chaining(self): |
| 196 | m = nn.Sequential(nn.Linear(2, 2), nn.ReLU(), nn.Linear(2, 1)) |
| 197 | pre_freeze_num_params = len(m.parameters()) |
| 198 | m.freeze().unfreeze() |
| 199 | self.assertEqual(len(m.parameters()), pre_freeze_num_params) |
| 200 | params_dict = m.parameters() |
| 201 | |
| 202 | self.assertFalse(m.update(params_dict).eval()._training) |
| 203 | self.assertTrue(m.train()._training) |
| 204 | |
| 205 | def test_quantize(self): |
| 206 | m = nn.Sequential(nn.Embedding(5, 256), nn.ReLU(), nn.Linear(256, 256)) |