diff --git a/lmdeploy/model.py b/lmdeploy/model.py index 284f266539..7d6981810c 100644 --- a/lmdeploy/model.py +++ b/lmdeploy/model.py @@ -92,16 +92,14 @@ def to_json(self, file_path=None): @classmethod def from_json(cls, file_or_string): """Construct a dataclass instance from a JSON file or JSON string.""" - try: - # Try to open the input_data as a file path + if os.path.isfile(file_or_string): with open(file_or_string, encoding='utf-8') as file: json_data = file.read() - except FileNotFoundError: - # If it's not a file path, assume it's a JSON string + else: + # If it's not a file path, assume it's a JSON string. Opening it + # as a path may fail with errors other than FileNotFoundError, + # e.g. a name longer than the OS limit, or quotes on Windows. json_data = file_or_string - except OSError: - # If it's not a file path and not a valid JSON string, raise error - raise ValueError('Invalid input. Must be a file path or a valid JSON string.') json_data = json.loads(json_data) if json_data.get('model_name', None) is None: json_data['model_name'] = random_uuid() diff --git a/tests/test_lmdeploy/test_model.py b/tests/test_lmdeploy/test_model.py index e9e4a32c10..3356abda58 100644 --- a/tests/test_lmdeploy/test_model.py +++ b/tests/test_lmdeploy/test_model.py @@ -87,6 +87,21 @@ def test_base_model(): assert model.messages2prompt('test') == 'test' +def test_chat_template_config_from_json(tmp_path): + import json + + from lmdeploy.model import ChatTemplateConfig + + # A long meta instruction makes the string invalid as a file name. + config = dict(model_name='from-json-test', meta_instruction='You are a helpful assistant. ' * 12) + json_str = json.dumps(config) + assert ChatTemplateConfig.from_json(json_str).meta_instruction == config['meta_instruction'] + + json_file = tmp_path / 'chat_template.json' + json_file.write_text(json_str, encoding='utf-8') + assert ChatTemplateConfig.from_json(str(json_file)).meta_instruction == config['meta_instruction'] + + def test_vicuna(): prompt = 'hello, can u introduce yourself' model = MODELS.get('vicuna')(capability='completion')