# # Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. # from types import SimpleNamespace import torch import vllm_ascend.model_loader.rfork.transfer_backend as transfer_backend from vllm_ascend.model_loader.rfork.transfer_backend import ( RForkTransferBackend, _parse_weight_info, _reshape_tensor_to_seed_shape, get_remote_instance_transfer_engine_info, ) def test_parse_weight_info_keeps_backward_compatibility(): assert _parse_weight_info([1, 2, 4]) == (1, 2, 4, None) def test_parse_weight_info_accepts_shape_metadata_from_json(): assert _parse_weight_info([1, 6, 2, [2, 3]]) == (1, 6, 2, (2, 3)) def test_parse_weight_info_rejects_invalid_shape_metadata(): assert _parse_weight_info([1, 6, 2, ["2", 3]]) is None assert _parse_weight_info([1, 6, 2, -1]) is None def test_reshape_tensor_to_seed_shape_updates_tensor_metadata_only(): tensor = torch.arange(6).reshape(2, 3) original_ptr = tensor.data_ptr() assert _reshape_tensor_to_seed_shape("weight", tensor, (1, 2, 3)) assert tuple(tensor.shape) == (1, 2, 3) assert tensor.data_ptr() == original_ptr def test_reshape_tensor_to_seed_shape_rejects_numel_mismatch(): tensor = torch.arange(6).reshape(2, 3) assert not _reshape_tensor_to_seed_shape("weight", tensor, (2, 2)) assert tuple(tensor.shape) == (2, 3) def test_recv_from_source_refreshes_registered_shape_after_reshape(monkeypatch): tensor = torch.arange(6).reshape(2, 3) backend = RForkTransferBackend.__new__(RForkTransferBackend) backend.rfork_transfer_engine = SimpleNamespace( batch_transfer_sync_read=lambda *args: SimpleNamespace(is_error=lambda: False) ) backend.rfork_transfer_engine_weights_shape_dict = {"weight": (2, 3)} monkeypatch.setattr(transfer_backend, "_iter_transferable_tensors", lambda model: iter([("weight", tensor)])) monkeypatch.setattr( transfer_backend, "get_remote_instance_transfer_engine_info", lambda *args: ( "seed-session", {"weight": [1, tensor.numel(), tensor.element_size()]}, {"weight": [1, 2, 3]}, ), ) assert backend.recv_from_source(object(), "127.0.0.1", 8000, "seed-key") assert tuple(tensor.shape) == (1, 2, 3) assert backend.rfork_transfer_engine_weights_shape_dict["weight"] == (1, 2, 3) def test_recv_from_source_reuses_registered_transferable_tensors(monkeypatch): tensor = torch.arange(6).reshape(2, 3) backend = RForkTransferBackend.__new__(RForkTransferBackend) backend.rfork_transfer_engine = SimpleNamespace( batch_transfer_sync_read=lambda *args: SimpleNamespace(is_error=lambda: False) ) backend.rfork_transfer_engine_weights_shape_dict = {"weight": (2, 3)} backend._registered_transferable_tensors = [("weight", tensor)] def fail_if_rescanned(model): raise AssertionError("recv_from_source should reuse the registered tensor cache") monkeypatch.setattr(transfer_backend, "_iter_transferable_tensors", fail_if_rescanned) monkeypatch.setattr( transfer_backend, "get_remote_instance_transfer_engine_info", lambda *args: ( "seed-session", {"weight": [1, tensor.numel(), tensor.element_size(), [2, 3]]}, None, ), ) assert backend.recv_from_source(object(), "127.0.0.1", 8000, "seed-key") assert backend._registered_transferable_tensors is None def test_get_remote_instance_transfer_engine_info_non_200_returns_three_values(monkeypatch): monkeypatch.setattr( transfer_backend.requests, "get", lambda *args, **kwargs: SimpleNamespace(status_code=503), ) assert get_remote_instance_transfer_engine_info("http://seed", "seed-key") == (None, None, None)