已合并
test(distributed):add test for torch.distributed.elastic.metrics.api.ConsoleMetricHandler, torch.distributed.elastic.metrics.api.NullMetricHandler and torch.distributed.elastic.metrics.configure #34726
test(distributed):add test for torch.distributed.elastic.metrics.api.ConsoleMetricHandler, torch.distributed.elastic.metrics.api.NullMetricHandler and torch.distributed.elastic.metrics.configure #34726
已合并
m0_73361278创建于 4月29日
共 1 个文件变更+149-0
@@ -0,0 +1,149 @@
1+"""
2+Add validation cases for torch.distributed.elastic.metrics APIs:
3+ 
4+1. PyTorch community lacks sufficient API validations, so this file is added.
5+2. This file validates the following apis:
6+torch.distributed.elastic.metrics.api.ConsoleMetricHandler
7+torch.distributed.elastic.metrics.api.NullMetricHandler
8+torch.distributed.elastic.metrics.configure
9+(extendable)
10+"""
11+ 
12+import unittest.mock as mock
13+ 
14+import torch.distributed.elastic.metrics.api as metrics_api
15+from torch.testing._internal.common_utils import TestCase, run_tests
16+ 
17+from torch.distributed.elastic.metrics import configure
18+from torch.distributed.elastic.metrics.api import (
19+ ConsoleMetricHandler,
20+ MetricData,
21+ MetricHandler,
22+ NullMetricHandler,
23+)
24+ 
25+ 
26+class ElasticMetricsApiTest(TestCase):
27+ 
28+ def setUp(self):
29+ super().setUp()
30+ self._original_default_metrics_handler = metrics_api._default_metrics_handler
31+ self._original_metrics_map = metrics_api._metrics_map.copy()
32+ self.addCleanup(self._restore_metrics_state)
33+ metrics_api._metrics_map.clear()
34+ 
35+ def _restore_metrics_state(self):
36+ metrics_api._default_metrics_handler = self._original_default_metrics_handler
37+ metrics_api._metrics_map.clear()
38+ metrics_api._metrics_map.update(self._original_metrics_map)
39+ 
40+ def test_console_metric_handler_emit(self):
41+ handler = ConsoleMetricHandler()
42+ metric_data = MetricData(
43+ timestamp=1,
44+ group_name="test_group",
45+ name="test_metric",
46+ value=1,
47+ )
48+ 
49+ with mock.patch("builtins.print") as print_mock:
50+ ret = handler.emit(metric_data)
51+ 
52+ self.assertIsNone(ret)
53+ self.assertIsInstance(handler, MetricHandler)
54+ self.assertEqual(1, print_mock.call_count)
55+ 
56+ output = print_mock.call_args[0][0]
57+ self.assertIn("1", output)
58+ self.assertIn("test_group", output)
59+ self.assertIn("test_metric", output)
60+ 
61+ def test_null_metric_handler_emit(self):
62+ handler = NullMetricHandler()
63+ metric_data = MetricData(
64+ timestamp=1,
65+ group_name="test_group",
66+ name="test_metric",
67+ value=1,
68+ )
69+ 
70+ with mock.patch("builtins.print") as print_mock:
71+ ret = handler.emit(metric_data)
72+ 
73+ self.assertIsNone(ret)
74+ self.assertIsInstance(handler, MetricHandler)
75+ self.assertEqual(0, print_mock.call_count)
76+ 
77+ def test_configure_default_console_metric_handler(self):
78+ handler = ConsoleMetricHandler()
79+ 
80+ with mock.patch("builtins.print") as print_mock:
81+ configure(handler)
82+ stream = metrics_api.getStream("default_group")
83+ stream.add_value("default_metric", 1)
84+ 
85+ self.assertEqual(1, print_mock.call_count)
86+ 
87+ output = print_mock.call_args[0][0]
88+ self.assertIn("default_group", output)
89+ self.assertIn("default_metric", output)
90+ self.assertIn("1", output)
91+ 
92+ def test_configure_default_null_metric_handler(self):
93+ handler = NullMetricHandler()
94+ 
95+ with mock.patch("builtins.print") as print_mock:
96+ configure(handler)
97+ stream = metrics_api.getStream("default_group")
98+ ret = stream.add_value("default_metric", 1)
99+ 
100+ self.assertIsNone(ret)
101+ self.assertEqual(0, print_mock.call_count)
102+ 
103+ def test_configure_group_specific_console_metric_handler(self):
104+ default_handler = NullMetricHandler()
105+ group_handler = ConsoleMetricHandler()
106+ 
107+ with mock.patch("builtins.print") as print_mock:
108+ configure(default_handler)
109+ configure(group_handler, group="console_group")
110+ 
111+ group_stream = metrics_api.getStream("console_group")
112+ default_stream = metrics_api.getStream("default_group")
113+ 
114+ group_stream.add_value("group_metric", 2)
115+ default_stream.add_value("default_metric", 3)
116+ 
117+ self.assertEqual(1, print_mock.call_count)
118+ 
119+ output = print_mock.call_args[0][0]
120+ self.assertIn("console_group", output)
121+ self.assertIn("group_metric", output)
122+ self.assertIn("2", output)
123+ self.assertNotIn("default_metric", output)
124+ 
125+ def test_configure_group_specific_null_metric_handler(self):
126+ default_handler = ConsoleMetricHandler()
127+ group_handler = NullMetricHandler()
128+ 
129+ with mock.patch("builtins.print") as print_mock:
130+ configure(default_handler)
131+ configure(group_handler, group="null_group")
132+ 
133+ null_stream = metrics_api.getStream("null_group")
134+ default_stream = metrics_api.getStream("default_group")
135+ 
136+ null_stream.add_value("null_metric", 1)
137+ default_stream.add_value("default_metric", 2)
138+ 
139+ self.assertEqual(1, print_mock.call_count)
140+ 
141+ output = print_mock.call_args[0][0]
142+ self.assertIn("default_group", output)
143+ self.assertIn("default_metric", output)
144+ self.assertIn("2", output)
145+ self.assertNotIn("null_metric", output)
146+ 
147+ 
148+if __name__ == "__main__":
149+ run_tests()