已合并
test(utils): add data worker API coverage tests on NPU #37243
Jinfan Liu创建于 5月30日
test(utils): add data worker API coverage tests on NPU #37243
已合并
共 1 个文件变更+62-0
| @@ -0,0 +1,62 @@ | |||
| 1 | +""" | ||
| 2 | +Add validation cases for torch.utils.data worker/control APIs: | ||
| 3 | +1. PyTorch community tests cover several worker/control APIs through DataLoader call chains, so this file adds direct API validations. | ||
| 4 | +2. This file validates _MultiProcessingDataLoaderIter, _InfiniteConstantSampler, ManagerWatchdog, _IterableDatasetStopIteration, _ResumeIteration, _set_worker_signal_handlers, and _remove_worker_pids (extendable). | ||
| 5 | +""" | ||
| 6 | + | ||
| 7 | +import os | ||
| 8 | + | ||
| 9 | +from torch.testing._internal.common_utils import TestCase, run_tests | ||
| 10 | +from torch.utils.data import DataLoader, Dataset, dataloader | ||
| 11 | +from torch.utils.data._utils import signal_handling, worker | ||
| 12 | + | ||
| 13 | + | ||
| 14 | +class _TinyMapDataset(Dataset): | ||
| 15 | + | ||
| 16 | + def __len__(self): | ||
| 17 | + return 4 | ||
| 18 | + | ||
| 19 | + def __getitem__(self, index): | ||
| 20 | + return f"sample-{index}" | ||
| 21 | + | ||
| 22 | + | ||
| 23 | +class TestUtilsDataWorkerControlAPIs(TestCase): | ||
| 24 | + | ||
| 25 | + def test_infinite_constant_sampler_yields_none(self): | ||
| 26 | + sampler_iter = iter(dataloader._InfiniteConstantSampler()) | ||
| 27 | + | ||
| 28 | + self.assertEqual([next(sampler_iter) for _ in range(3)], [None, None, None]) | ||
| 29 | + | ||
| 30 | + def test_worker_control_message_fields(self): | ||
| 31 | + stop_message = worker._IterableDatasetStopIteration(worker_id=2) | ||
| 32 | + resume_message = worker._ResumeIteration(seed=123) | ||
| 33 | + | ||
| 34 | + self.assertEqual(stop_message.worker_id, 2) | ||
| 35 | + self.assertEqual(resume_message.seed, 123) | ||
| 36 | + self.assertIn("worker_id=2", repr(stop_message)) | ||
| 37 | + self.assertIn("seed=123", repr(resume_message)) | ||
| 38 | + | ||
| 39 | + def test_manager_watchdog_reports_parent_alive(self): | ||
| 40 | + watchdog = worker.ManagerWatchdog() | ||
| 41 | + | ||
| 42 | + self.assertTrue(watchdog.is_alive()) | ||
| 43 | + | ||
| 44 | + def test_worker_signal_handlers_and_pid_cleanup(self): | ||
| 45 | + loader_id = id(self) | ||
| 46 | + | ||
| 47 | + self.assertIsNone(signal_handling._set_worker_signal_handlers()) | ||
| 48 | + self.assertIsNone(signal_handling._set_worker_pids(loader_id, (os.getpid(),))) | ||
| 49 | + self.assertIsNone(signal_handling._remove_worker_pids(loader_id)) | ||
| 50 | + | ||
| 51 | + def test_multiprocessing_dataloader_iter_type_and_shutdown(self): | ||
| 52 | + loader = DataLoader(_TinyMapDataset(), batch_size=2, num_workers=1) | ||
| 53 | + iterator = iter(loader) | ||
| 54 | + try: | ||
| 55 | + self.assertIsInstance(iterator, dataloader._MultiProcessingDataLoaderIter) | ||
| 56 | + self.assertEqual(next(iterator), ["sample-0", "sample-1"]) | ||
| 57 | + finally: | ||
| 58 | + iterator._shutdown_workers() | ||
| 59 | + | ||
| 60 | + | ||
| 61 | +if __name__ == "__main__": | ||
| 62 | + run_tests() | ||