已合并
test(fx):Add validation cases for torch._C._distributed_rpc._is_current_rpc_agent_set on NPU #43172
test(fx):Add validation cases for torch._C._distributed_rpc._is_current_rpc_agent_set on NPU #43172
已合并
olpk创建于 7月29日
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()