MCPcopy Create free account
hub / github.com/microsoft/TRELLIS / from_pretrained

Function from_pretrained

trellis/models/__init__.py:40–70  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

38
39
40def 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

Callers

nothing calls this directly

Calls 3

__getattr__Function · 0.70
loadMethod · 0.45
load_state_dictMethod · 0.45

Tested by

no test coverage detected