已合并
test(fx):Add validation cases for torch._C._distributed_rpc._is_current_rpc_agent_set on NPU #43171
olpk创建于 28 天前
test(fx):Add validation cases for torch._C._distributed_rpc._is_current_rpc_agent_set on NPU #43171
已合并
共 1 个文件变更+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()) | ||
| 56 | + | ||
| 57 | + | ||
| 58 | +if __name__ == "__main__": | ||
| 59 | + run_tests() | ||
🟡 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。