已合并
fix: correct init_step_count and erase_step_count cache behavior #42170
fix: correct init_step_count and erase_step_count cache behavior #42170
已合并
chenlan114514创建于 7月20日
1 个文件变更+129-0
@@ -0,0 +1,129 @@
1+# Copyright (c) 2026 Huawei Technologies Co., Ltd. All rights reserved.
2+#
3+# This source code is licensed under the BSD-style license found in the
4+# LICENSE file in the root directory of this source tree.
5+ 
6+from torch.testing._internal.common_utils import TestCase, run_tests
7+from torch.autograd.profiler import KinetoStepTracker
8+ 
9+ 
10+class TestKinetoStepTracker(TestCase):
11+ """Unit tests for KinetoStepTracker native behavior contract.
12+ 
13+ Verifies that the upstream KinetoStepTracker implementation behaves
14+ according to its design contract:
15+ - Global step is monotonically non-decreasing.
16+ - init_step_count aligns new requester to current global step.
17+ - increment_step advances per-requester step and updates global max.
18+ - erase_step_count removes a requester but never rolls back global step.
19+ - erase_step_count returns bool indicating deletion success.
20+ """
21+ 
22+ def setUp(self):

setUp() 直接清空 KinetoStepTracker._step_dict 并设置 _current_step = 0,但没有 tearDown() 快照并恢复测试前状态。

likedislike
chenlan114514
chenlan114514
7月30日 评论:
23+ # Save snapshot of original global state for post-test restoration.
24+ self._saved_step_dict = dict(KinetoStepTracker._step_dict)
25+ self._saved_current_step = KinetoStepTracker._current_step
26+ 
27+ # Reset to clean initial state to ensure deterministic test results.
28+ KinetoStepTracker._step_dict.clear()
29+ KinetoStepTracker._current_step = 0
30+ 
31+ def tearDown(self):
32+ # Fully restore original global state to prevent cross-test pollution.
33+ KinetoStepTracker._step_dict.clear()
34+ KinetoStepTracker._step_dict.update(self._saved_step_dict)
35+ KinetoStepTracker._current_step = self._saved_current_step
36+ 
37+ def test_init_does_not_alter_global_step(self):
38+ """init_step_count only registers requester, never changes global step."""
39+ self.assertEqual(KinetoStepTracker.current_step(), 0)
40+ 
41+ KinetoStepTracker.init_step_count("req1")
42+ self.assertEqual(KinetoStepTracker.current_step(), 0)
43+ 
44+ # Global step advances only through increment_step
45+ KinetoStepTracker.increment_step("req1")
46+ self.assertEqual(KinetoStepTracker.current_step(), 1)
47+ 
48+ # New requester aligns to current step; global max remains unchanged
49+ KinetoStepTracker.init_step_count("req2")
50+ self.assertEqual(KinetoStepTracker.current_step(), 1)
51+ self.assertEqual(KinetoStepTracker._step_dict.get("req2"), 1)
52+ 
53+ def test_increment_takes_maximum(self):
54+ """current_step always equals the maximum step among all requesters."""
55+ KinetoStepTracker.init_step_count("req1")
56+ KinetoStepTracker.init_step_count("req2")
57+ self.assertEqual(KinetoStepTracker.current_step(), 0)
58+ 
59+ KinetoStepTracker.increment_step("req1")
60+ self.assertEqual(KinetoStepTracker.current_step(), 1)
61+ 
62+ KinetoStepTracker.increment_step("req2")
63+ KinetoStepTracker.increment_step("req2")
64+ self.assertEqual(KinetoStepTracker.current_step(), 2)
65+ 
66+ def test_erase_keeps_step_monotonic(self):
67+ """Erasing any requester never decreases global step (monotonic contract)."""
68+ KinetoStepTracker.init_step_count("req1")
69+ KinetoStepTracker.init_step_count("req2")
70+ 
71+ KinetoStepTracker.increment_step("req1")
72+ KinetoStepTracker.increment_step("req1") # req1=2, req2=0
73+ self.assertEqual(KinetoStepTracker.current_step(), 2)
74+ 
75+ # Erasing the max-step requester does not roll back global step
76+ KinetoStepTracker.erase_step_count("req1")
77+ self.assertEqual(KinetoStepTracker.current_step(), 2)
78+ 
79+ def test_erase_all_requesters_keeps_history(self):
80+ """Global step retains historical maximum even after all requesters are erased."""
81+ KinetoStepTracker.init_step_count("req1")
82+ for _ in range(3):
83+ KinetoStepTracker.increment_step("req1")
84+ self.assertEqual(KinetoStepTracker.current_step(), 3)
85+ 
86+ KinetoStepTracker.erase_step_count("req1")
87+ # Step does not roll back even with no active requester
88+ self.assertEqual(KinetoStepTracker.current_step(), 3)
89+ 
90+ def test_erase_return_value_contract(self):
91+ """erase_step_count returns bool indicating whether deletion succeeded."""
92+ KinetoStepTracker.init_step_count("req1")
93+ 
94+ # Returns True for an existing requester
95+ self.assertTrue(KinetoStepTracker.erase_step_count("req1"))
96+ # Returns False for an already-deleted requester
97+ self.assertFalse(KinetoStepTracker.erase_step_count("req1"))
98+ # Returns False for a non-existent requester
99+ self.assertFalse(KinetoStepTracker.erase_step_count("nonexistent"))
100+ 
101+ def test_reinit_is_idempotent(self):
102+ """Calling init_step_count repeatedly on the same requester has no side effect."""
103+ KinetoStepTracker.init_step_count("req1")
104+ KinetoStepTracker.increment_step("req1")
105+ self.assertEqual(KinetoStepTracker.current_step(), 1)
106+ self.assertEqual(KinetoStepTracker._step_dict.get("req1"), 1)
107+ 
108+ # Re-initializing an existing requester does not change its step count
109+ KinetoStepTracker.init_step_count("req1")
110+ self.assertEqual(KinetoStepTracker.current_step(), 1)
111+ self.assertEqual(KinetoStepTracker._step_dict.get("req1"), 1)
112+ 
113+ def test_new_requester_inherits_current_step(self):
114+ """New requester inherits current global step for alignment."""
115+ KinetoStepTracker.init_step_count("req1")
116+ for _ in range(5):
117+ KinetoStepTracker.increment_step("req1")
118+ self.assertEqual(KinetoStepTracker.current_step(), 5)
119+ 
120+ KinetoStepTracker.erase_step_count("req1")
121+ KinetoStepTracker.init_step_count("req2")
122+ # Global step stays at historical maximum
123+ self.assertEqual(KinetoStepTracker.current_step(), 5)
124+ # New requester aligns to the current global step
125+ self.assertEqual(KinetoStepTracker._step_dict.get("req2"), 5)
126+ 
127+ 
128+if __name__ == "__main__":
129+ run_tests()