[Bugfix] Fix Mistral v0.3 Weight Loading (#5005)

Co-authored-by: Cody Yu <hao.yu.cody@gmail.com>
This commit is contained in:
Robert Shaw
2024-05-24 14:28:27 +02:00
committed by GitHub
parent 6a50f4cafa
commit 919770957f
3 changed files with 79 additions and 3 deletions

View File

@@ -23,7 +23,8 @@ from vllm.model_executor.model_loader.tensorizer import (
from vllm.model_executor.model_loader.utils import (get_model_architecture,
set_default_torch_dtype)
from vllm.model_executor.model_loader.weight_utils import (
download_weights_from_hf, filter_files_not_needed_for_inference,
download_safetensors_index_file_from_hf, download_weights_from_hf,
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
get_quant_config, initialize_dummy_weights, np_cache_weights_iterator,
pt_weights_iterator, safetensors_weights_iterator)
from vllm.model_executor.models.vlm_base import VisionLanguageModelBase
@@ -188,7 +189,19 @@ class DefaultModelLoader(BaseModelLoader):
use_safetensors = True
break
if not use_safetensors:
if use_safetensors:
# For models like Mistral-7B-Instruct-v0.3
# there are both sharded safetensors files and a consolidated
# safetensors file. Using both breaks.
# Here, we download the `model.safetensors.index.json` and filter
# any files not found in the index.
if not is_local:
download_safetensors_index_file_from_hf(
model_name_or_path, self.load_config.download_dir,
revision)
hf_weights_files = filter_duplicate_safetensors_files(
hf_weights_files, hf_folder)
else:
hf_weights_files = filter_files_not_needed_for_inference(
hf_weights_files)