Loading tests/test_train.py +5 −0 Changes for tests/test_train.py: 5 added lines, 0 removed lines. Original line number Diff line number Diff line Loading @@ -5,6 +5,8 @@ # the root directory of this source tree. An additional grant of patent rights # can be found in the PATENTS file in the same directory. import contextlib from io import StringIO import unittest from unittest.mock import MagicMock, patch Loading Loading @@ -37,6 +39,7 @@ class TestLoadCheckpoint(unittest.TestCase): [p.start() for p in self.applied_patches] def test_load_partial_checkpoint(self): with contextlib.redirect_stdout(StringIO()): trainer = mock_trainer(2, 200, False) loader = mock_loader(150) epoch, ds = train.load_checkpoint(MagicMock(), trainer, loader) Loading @@ -44,6 +47,7 @@ class TestLoadCheckpoint(unittest.TestCase): self.assertEqual(next(ds), 50) def test_load_full_checkpoint(self): with contextlib.redirect_stdout(StringIO()): trainer = mock_trainer(2, 300, True) loader = mock_loader(150) epoch, ds = train.load_checkpoint(MagicMock(), trainer, loader) Loading @@ -51,6 +55,7 @@ class TestLoadCheckpoint(unittest.TestCase): self.assertEqual(next(iter(ds)), 0) def test_load_no_checkpoint(self): with contextlib.redirect_stdout(StringIO()): trainer = mock_trainer(0, 0, False) loader = mock_loader(150) self.patches['os.path.isfile'].return_value = False Loading Loading
tests/test_train.py +5 −0 Changes for tests/test_train.py: 5 added lines, 0 removed lines. Original line number Diff line number Diff line Loading @@ -5,6 +5,8 @@ # the root directory of this source tree. An additional grant of patent rights # can be found in the PATENTS file in the same directory. import contextlib from io import StringIO import unittest from unittest.mock import MagicMock, patch Loading Loading @@ -37,6 +39,7 @@ class TestLoadCheckpoint(unittest.TestCase): [p.start() for p in self.applied_patches] def test_load_partial_checkpoint(self): with contextlib.redirect_stdout(StringIO()): trainer = mock_trainer(2, 200, False) loader = mock_loader(150) epoch, ds = train.load_checkpoint(MagicMock(), trainer, loader) Loading @@ -44,6 +47,7 @@ class TestLoadCheckpoint(unittest.TestCase): self.assertEqual(next(ds), 50) def test_load_full_checkpoint(self): with contextlib.redirect_stdout(StringIO()): trainer = mock_trainer(2, 300, True) loader = mock_loader(150) epoch, ds = train.load_checkpoint(MagicMock(), trainer, loader) Loading @@ -51,6 +55,7 @@ class TestLoadCheckpoint(unittest.TestCase): self.assertEqual(next(iter(ds)), 0) def test_load_no_checkpoint(self): with contextlib.redirect_stdout(StringIO()): trainer = mock_trainer(0, 0, False) loader = mock_loader(150) self.patches['os.path.isfile'].return_value = False Loading