119 lines
4.3 KiB
Python
119 lines
4.3 KiB
Python
#
|
|
# 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)
|