diff --git a/tests/model_executor/test_cohere_asr.py b/tests/model_executor/test_cohere_asr.py new file mode 100644 index 000000000000..fcd1bdc9874a --- /dev/null +++ b/tests/model_executor/test_cohere_asr.py @@ -0,0 +1,51 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import torch +from torch import nn + +from vllm.model_executor.models.cohere_asr import ( + CohereASRModel, + RelPositionMultiHeadAttention, +) +from vllm.utils.torch_utils import set_default_torch_dtype + + +class _CohereASRWeightLoadingModel(CohereASRModel): + def __init__(self) -> None: + nn.Module.__init__(self) + self.register_buffer("pos_bias_u", torch.zeros(2, 2)) + self.register_buffer("pos_bias_v", torch.zeros(2, 2)) + self.attention = RelPositionMultiHeadAttention( + n_head=2, + n_feat=4, + pos_bias_u=self.pos_bias_u, + pos_bias_v=self.pos_bias_v, + ) + + +@pytest.mark.cpu_test +def test_load_weights_preserves_runtime_bias_dtype() -> None: + with set_default_torch_dtype(torch.float16): + model = _CohereASRWeightLoadingModel() + + loaded_weight = torch.ones_like(model.pos_bias_u, dtype=torch.float32) + model.load_weights([("pos_bias_u", loaded_weight), ("pos_bias_v", loaded_weight)]) + + assert model.pos_bias_u.dtype == torch.float16 + assert model.pos_bias_v.dtype == torch.float16 + assert model.attention.pos_bias_u.dtype == torch.float16 + assert model.attention.pos_bias_v.dtype == torch.float16 + torch.testing.assert_close(model.pos_bias_u, loaded_weight.half()) + + output = model.attention( + query=torch.randn(1, 3, 4, dtype=torch.float16), + key=torch.randn(1, 3, 4, dtype=torch.float16), + value=torch.randn(1, 3, 4, dtype=torch.float16), + mask=None, + pos_emb=torch.randn(1, 5, 4, dtype=torch.float16), + ) + + assert output.dtype == torch.float16 + assert output.shape == (1, 3, 4) diff --git a/vllm/model_executor/models/cohere_asr.py b/vllm/model_executor/models/cohere_asr.py index b4e72d5b3f68..ef99fb4a7dff 100644 --- a/vllm/model_executor/models/cohere_asr.py +++ b/vllm/model_executor/models/cohere_asr.py @@ -1817,16 +1817,6 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: param = params_dict[name] weight_loader = getattr(param, "weight_loader", default_weight_loader) - # Convert buffer dtype to match loaded weight for pos_bias tensors - if "pos_bias" in name and param.dtype != loaded_weight.dtype: - logger.info( - "Converting buffer %s dtype from %s to %s for loading.", - name, - param.dtype, - loaded_weight.dtype, - ) - param.data = param.data.to(loaded_weight.dtype) - weight_loader(param, loaded_weight) loaded_params.add(name) return loaded_params