已合并
test(utils): add data worker API coverage tests on NPU #37242
test(utils): add data worker API coverage tests on NPU #37242
已合并
Jinfan Liu创建于 5月30日
1 个文件变更+62-0
Atest/test_utils_data_api.py+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()