(self)
| 22 | |
| 23 | class TestSetDeterminism(unittest.TestCase): |
| 24 | def test_values(self): |
| 25 | # check system default flags |
| 26 | set_determinism(None) |
| 27 | self.assertTrue(not torch.backends.cudnn.deterministic) |
| 28 | self.assertTrue(get_seed() is None) |
| 29 | # set default seed |
| 30 | set_determinism() |
| 31 | self.assertTrue(get_seed() is not None) |
| 32 | self.assertTrue(torch.backends.cudnn.deterministic) |
| 33 | self.assertTrue(not torch.backends.cudnn.benchmark) |
| 34 | # resume default |
| 35 | set_determinism(None) |
| 36 | self.assertTrue(not torch.backends.cudnn.deterministic) |
| 37 | self.assertTrue(not torch.backends.cudnn.benchmark) |
| 38 | self.assertTrue(get_seed() is None) |
| 39 | # test seeds |
| 40 | seed = 255 |
| 41 | set_determinism(seed=seed) |
| 42 | self.assertEqual(seed, get_seed()) |
| 43 | a = np.random.randint(seed) |
| 44 | b = torch.randint(seed, (1,)) |
| 45 | |
| 46 | # test when global flag support is disabled |
| 47 | torch.backends.disable_global_flags() |
| 48 | set_determinism(seed=seed) |
| 49 | c = np.random.randint(seed) |
| 50 | d = torch.randint(seed, (1,)) |
| 51 | self.assertEqual(a, c) |
| 52 | self.assertEqual(b, d) |
| 53 | self.assertTrue(torch.backends.cudnn.deterministic) |
| 54 | self.assertTrue(not torch.backends.cudnn.benchmark) |
| 55 | set_determinism(seed=None) |
| 56 | |
| 57 | |
| 58 | class TestSetFlag(unittest.TestCase): |
nothing calls this directly
no test coverage detected