已合并
test(fx):Add validation cases for torch._C._distributed_rpc._is_current_rpc_agent_set on NPU #43171
test(fx):Add validation cases for torch._C._distributed_rpc._is_current_rpc_agent_set on NPU #43171
已合并
olpk创建于 28 天前
1 个文件变更+59-0
Atest/distributed/rpc/test_rpc_agent_set.py+59-0
@@ -0,0 +1,59 @@
1+# Copyright (c) 2026 Huawei Technologies Co., Ltd
2+# All rights reserved.
3+#
4+# Licensed under the BSD 3-Clause License (the "License");
5+# you may not use this file except in compliance with the License.
6+# You may obtain a copy of the License at
7+#
8+# https://opensource.org/licenses/BSD-3-Clause
9+#
10+# Unless required by applicable law or agreed to in writing, software
11+# distributed under the License is distributed on an "AS IS" BASIS,
12+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+# See the License for the specific language governing permissions and
14+# limitations under the License.
15+ 
16+"""
17+Add validation cases for torch._C._distributed_rpc APIs on NPU:
18+1. PyTorch community lacks direct and sufficient validation for some APIs,
19+ so this file is added.
20+2. This file validates torch._C._distributed_rpc._is_current_rpc_agent_set (extendable).
21+"""
22+ 
23+import torch # noqa: F401
24+from torch_npu.testing.testcase import run_tests, TestCase
25+ 
26+ 
27+class TestIsCurrentRpcAgentSet(TestCase):
28+ """Test cases for torch._C._distributed_rpc._is_current_rpc_agent_set."""
29+ 
30+ def test_is_current_rpc_agent_set_import(self):
31+ """Verify that _is_current_rpc_agent_set is importable and callable."""
32+ from torch._C._distributed_rpc import _is_current_rpc_agent_set
33+ self.assertTrue(callable(_is_current_rpc_agent_set))
34+ 
35+ def test_is_current_rpc_agent_set_default(self):
36+ """Verify that _is_current_rpc_agent_set returns False when RPC is not initialized."""
37+ from torch._C._distributed_rpc import _is_current_rpc_agent_set
38+ self.assertFalse(_is_current_rpc_agent_set())
39+ 
40+ def test_is_current_rpc_agent_set_after_init(self):
41+ """Verify that _is_current_rpc_agent_set returns True after init_rpc, and False after shutdown."""
42+ import os
43+ import torch.distributed.rpc as rpc
44+ from torch._C._distributed_rpc import _is_current_rpc_agent_set
45+ 
46+ os.environ.setdefault("MASTER_ADDR", "localhost")
47+ os.environ.setdefault("MASTER_PORT", "29500")
48+ 
49+ self.assertFalse(_is_current_rpc_agent_set())
50+ rpc.init_rpc("worker0", rank=0, world_size=1)
51+ try:
52+ self.assertTrue(_is_current_rpc_agent_set())
53+ finally:
54+ rpc.shutdown()
55+ self.assertFalse(_is_current_rpc_agent_set())
atomgit-bot
atomgit-botatomgit-bot28 天前

🟡 Medium Priority

test_is_current_rpc_agent_set_after_init 方法中,第 46 行 rpc.init_rpc(...) 成功后,若第 47 行 self.assertTrue(...) 断言失败,异常会直接向上传播,第 48 行 rpc.shutdown() 永远无法执行。此时 RPC agent 仍处于已初始化状态,同一进程内的后续测试(如 test_is_current_rpc_agent_set_default)调用 _is_current_rpc_agent_set() 将返回 True 而非预期的 False,导致级联测试失败或产生误导性结果。

触发条件:init_rpc 成功但 RPC agent 内部状态异常(极端情况)导致 _is_current_rpc_agent_set() 返回了意外值。

修复方向:使用 try/finally 确保 rpc.shutdown() 始终被执行;或通过 setUp/tearDown 管理 RPC 生命周期。

建议:使用 try/finally 包裹 init_rpc 之后的核心断言,确保无论断言成功与否 shutdown 都被调用:将第 46-49 行改为 rpc.init_rpc(...)try: self.assertTrue(...)finally: rpc.shutdown() → 在 finally 后(或内部)验证 shutdown 后返回 False。

改动建议
55
+ def test_is_current_rpc_agent_set_after_init(self):
56
+ """Verify that _is_current_rpc_agent_set returns True after init_rpc, and False after shutdown."""
57
+ import torch.distributed.rpc as rpc
58
+ from torch._C._distributed_rpc import _is_current_rpc_agent_set
59
+
60
+ self.assertFalse(_is_current_rpc_agent_set())
61
+ rpc.init_rpc("worker0", rank=0, world_size=1)
62
+ try:
63
+ self.assertTrue(_is_current_rpc_agent_set())
64
+ finally:
65
+ rpc.shutdown()
55
66
  self.assertFalse(_is_current_rpc_agent_set())
应用建议
likedislike
群青世界
16 天前 评论:
56+ 
57+ 
58+if __name__ == "__main__":
59+ run_tests()