Load a model from a pretrained checkpoint. Args: path: The path to the checkpoint. Can be either local path or a Hugging Face model name. NOTE: config file and model file should take the name f'{path}.json' and f'{path}.safetensors' respectively. **kwargs: Add
(path: str, **kwargs)
| 38 | |
| 39 | |
| 40 | def from_pretrained(path: str, **kwargs): |
| 41 | """ |
| 42 | Load a model from a pretrained checkpoint. |
| 43 | |
| 44 | Args: |
| 45 | path: The path to the checkpoint. Can be either local path or a Hugging Face model name. |
| 46 | NOTE: config file and model file should take the name f'{path}.json' and f'{path}.safetensors' respectively. |
| 47 | **kwargs: Additional arguments for the model constructor. |
| 48 | """ |
| 49 | import os |
| 50 | import json |
| 51 | from safetensors.torch import load_file |
| 52 | is_local = os.path.exists(f"{path}.json") and os.path.exists(f"{path}.safetensors") |
| 53 | |
| 54 | if is_local: |
| 55 | config_file = f"{path}.json" |
| 56 | model_file = f"{path}.safetensors" |
| 57 | else: |
| 58 | from huggingface_hub import hf_hub_download |
| 59 | path_parts = path.split('/') |
| 60 | repo_id = f'{path_parts[0]}/{path_parts[1]}' |
| 61 | model_name = '/'.join(path_parts[2:]) |
| 62 | config_file = hf_hub_download(repo_id, f"{model_name}.json") |
| 63 | model_file = hf_hub_download(repo_id, f"{model_name}.safetensors") |
| 64 | |
| 65 | with open(config_file, 'r') as f: |
| 66 | config = json.load(f) |
| 67 | model = __getattr__(config['name'])(**config['args'], **kwargs) |
| 68 | model.load_state_dict(load_file(model_file)) |
| 69 | |
| 70 | return model |
| 71 | |
| 72 | |
| 73 | # For Pylance |
nothing calls this directly
no test coverage detected