已合并
fix: correct init_step_count and erase_step_count cache behavior #42170
chenlan114514创建于 7月20日
fix: correct init_step_count and erase_step_count cache behavior #42170
已合并
共 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): | ||
群 | |||
| 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() | ||
setUp()直接清空KinetoStepTracker._step_dict并设置_current_step = 0,但没有tearDown()快照并恢复测试前状态。