codekingpro/portable-devtools
114k
1import inspect2import json3import os4from collections.abc import Callable5from dataclasses import Field, asdict, dataclass, is_dataclass6from pathlib import Path7from typing import Any, ClassVar, Protocol, TypeVar8 9import packaging.version10 11from . import constants12from .errors import EntryNotFoundError, HfHubHTTPError13from .file_download import hf_hub_download14from .hf_api import HfApi15from .repocard import ModelCard, ModelCardData16from .utils import (17 SoftTemporaryDirectory,18 is_jsonable,19 is_safetensors_available,20 is_simple_optional_type,21 is_torch_available,22 logging,23 unwrap_simple_optional_type,24 validate_hf_hub_args,25)26 27 28if is_torch_available():29 import torch # type: ignore30 31if is_safetensors_available():32 import safetensors33 from safetensors.torch import load_model as load_model_as_safetensor34 from safetensors.torch import save_model as save_model_as_safetensor35 36 37logger = logging.get_logger(__name__)38 39 40# Type alias for dataclass instances, copied from https://github.com/python/typeshed/blob/9f28171658b9ca6c32a7cb93fbb99fc92b17858b/stdlib/_typeshed/__init__.pyi#L34941class DataclassInstance(Protocol):42 __dataclass_fields__: ClassVar[dict[str, Field]]43 44 45# Generic variable that is either ModelHubMixin or a subclass thereof46T = TypeVar("T", bound="ModelHubMixin")47# Generic variable to represent an args type48ARGS_T = TypeVar("ARGS_T")49ENCODER_T = Callable[[ARGS_T], Any]50DECODER_T = Callable[[Any], ARGS_T]51CODER_T = tuple[ENCODER_T, DECODER_T]52 53 54DEFAULT_MODEL_CARD = """55---56# For reference on model card metadata, see the spec: https://github.com/huggingface/hub-docs/blob/main/modelcard.md?plain=157# Doc / guide: https://huggingface.co/docs/hub/model-cards58{{ card_data }}59---60 61This model has been pushed to the Hub using the [PytorchModelHubMixin](https://huggingface.co/docs/huggingface_hub/package_reference/mixins#huggingface_hub.PyTorchModelHubMixin) integration:62- Code: {{ repo_url | default("[More Information Needed]", true) }}63- Paper: {{ paper_url | default("[More Information Needed]", true) }}64- Docs: {{ docs_url | default("[More Information Needed]", true) }}65"""66 67 68@dataclass69class MixinInfo:70 model_card_template: str71 model_card_data: ModelCardData72 docs_url: str | None = None73 paper_url: str | None = None74 repo_url: str | None = None75 76 77class ModelHubMixin:78 """79 A generic mixin to integrate ANY machine learning framework with the Hub.80 81 To integrate your framework, your model class must inherit from this class. Custom logic for saving/loading models82 have to be overwritten in [`_from_pretrained`] and [`_save_pretrained`]. [`PyTorchModelHubMixin`] is a good example83 of mixin integration with the Hub. Check out our [integration guide](../guides/integrations) for more instructions.84 85 When inheriting from [`ModelHubMixin`], you can define class-level attributes. These attributes are not passed to86 `__init__` but to the class definition itself. This is useful to define metadata about the library integrating87 [`ModelHubMixin`].88 89 For more details on how to integrate the mixin with your library, checkout the [integration guide](../guides/integrations).90 91 Args:92 repo_url (`str`, *optional*):93 URL of the library repository. Used to generate model card.94 paper_url (`str`, *optional*):95 URL of the library paper. Used to generate model card.96 docs_url (`str`, *optional*):97 URL of the library documentation. Used to generate model card.98 model_card_template (`str`, *optional*):99 Template of the model card. Used to generate model card. Defaults to a generic template.100 language (`str` or `list[str]`, *optional*):101 Language supported by the library. Used to generate model card.102 library_name (`str`, *optional*):103 Name of the library integrating ModelHubMixin. Used to generate model card.104 license (`str`, *optional*):105 License of the library integrating ModelHubMixin. Used to generate model card.106 E.g: "apache-2.0"107 license_name (`str`, *optional*):108 Name of the library integrating ModelHubMixin. Used to generate model card.109 Only used if `license` is set to `other`.110 E.g: "coqui-public-model-license".111 license_link (`str`, *optional*):112 URL to the license of the library integrating ModelHubMixin. Used to generate model card.113 Only used if `license` is set to `other` and `license_name` is set.114 E.g: "https://coqui.ai/cpml".115 pipeline_tag (`str`, *optional*):116 Tag of the pipeline. Used to generate model card. E.g. "text-classification".117 tags (`list[str]`, *optional*):118 Tags to be added to the model card. Used to generate model card. E.g. ["computer-vision"]119 coders (`dict[Type, tuple[Callable, Callable]]`, *optional*):120 Dictionary of custom types and their encoders/decoders. Used to encode/decode arguments that are not121 jsonable by default. E.g. dataclasses, argparse.Namespace, OmegaConf, etc.122 123 Example:124 125 ```python126 >>> from huggingface_hub import ModelHubMixin127 128 # Inherit from ModelHubMixin129 >>> class MyCustomModel(130 ... ModelHubMixin,131 ... library_name="my-library",132 ... tags=["computer-vision"],133 ... repo_url="https://github.com/huggingface/my-cool-library",134 ... paper_url="https://arxiv.org/abs/2304.12244",135 ... docs_url="https://huggingface.co/docs/my-cool-library",136 ... # ^ optional metadata to generate model card137 ... ):138 ... def __init__(self, size: int = 512, device: str = "cpu"):139 ... # define how to initialize your model140 ... super().__init__()141 ... ...142 ...143 ... def _save_pretrained(self, save_directory: Path) -> None:144 ... # define how to serialize your model145 ... ...146 ...147 ... @classmethod148 ... def from_pretrained(149 ... cls: type[T],150 ... pretrained_model_name_or_path: Union[str, Path],151 ... *,152 ... force_download: bool = False,153 ... token: Optional[Union[str, bool]] = None,154 ... cache_dir: Optional[Union[str, Path]] = None,155 ... local_files_only: bool = False,156 ... revision: Optional[str] = None,157 ... **model_kwargs,158 ... ) -> T:159 ... # define how to deserialize your model160 ... ...161 162 >>> model = MyCustomModel(size=256, device="gpu")163 164 # Save model weights to local directory165 >>> model.save_pretrained("my-awesome-model")166 167 # Push model weights to the Hub168 >>> model.push_to_hub("my-awesome-model")169 170 # Download and initialize weights from the Hub171 >>> reloaded_model = MyCustomModel.from_pretrained("username/my-awesome-model")172 >>> reloaded_model.size173 256174 175 # Model card has been correctly populated176 >>> from huggingface_hub import ModelCard177 >>> card = ModelCard.load("username/my-awesome-model")178 >>> card.data.tags179 ["x-custom-tag", "pytorch_model_hub_mixin", "model_hub_mixin"]180 >>> card.data.library_name181 "my-library"182 ```183 """184 185 _hub_mixin_config: dict | DataclassInstance | None = None186 # ^ optional config attribute automatically set in `from_pretrained`187 _hub_mixin_info: MixinInfo188 # ^ information about the library integrating ModelHubMixin (used to generate model card)189 _hub_mixin_inject_config: bool # whether `_from_pretrained` expects `config` or not190 _hub_mixin_init_parameters: dict[str, inspect.Parameter] # __init__ parameters191 _hub_mixin_jsonable_default_values: dict[str, Any] # default values for __init__ parameters192 _hub_mixin_jsonable_custom_types: tuple[type, ...] # custom types that can be encoded/decoded193 _hub_mixin_coders: dict[type, CODER_T] # encoders/decoders for custom types194 # ^ internal values to handle config195 196 def __init_subclass__(197 cls,198 *,199 # Generic info for model card200 repo_url: str | None = None,201 paper_url: str | None = None,202 docs_url: str | None = None,203 # Model card template204 model_card_template: str = DEFAULT_MODEL_CARD,205 # Model card metadata206 language: list[str] | None = None,207 library_name: str | None = None,208 license: str | None = None,209 license_name: str | None = None,210 license_link: str | None = None,211 pipeline_tag: str | None = None,212 tags: list[str] | None = None,213 # How to encode/decode arguments with custom type into a JSON config?214 coders: None215 | (216 dict[type, CODER_T]217 # Key is a type.218 # Value is a tuple (encoder, decoder).219 # Example: {MyCustomType: (lambda x: x.value, lambda data: MyCustomType(data))}220 ) = None,221 ) -> None:222 """Inspect __init__ signature only once when subclassing + handle modelcard."""223 super().__init_subclass__()224 225 # Will be reused when creating modelcard226 tags = tags or []227 tags.append("model_hub_mixin")228 229 # Initialize MixinInfo if not existent230 info = MixinInfo(model_card_template=model_card_template, model_card_data=ModelCardData())231 232 # If parent class has a MixinInfo, inherit from it as a copy233 if hasattr(cls, "_hub_mixin_info"):234 # Inherit model card template from parent class if not explicitly set235 if model_card_template == DEFAULT_MODEL_CARD:236 info.model_card_template = cls._hub_mixin_info.model_card_template237 238 # Inherit from parent model card data239 info.model_card_data = ModelCardData(**cls._hub_mixin_info.model_card_data.to_dict())240 241 # Inherit other info242 info.docs_url = cls._hub_mixin_info.docs_url243 info.paper_url = cls._hub_mixin_info.paper_url244 info.repo_url = cls._hub_mixin_info.repo_url245 cls._hub_mixin_info = info246 247 # Update MixinInfo with metadata248 if model_card_template is not None and model_card_template != DEFAULT_MODEL_CARD:249 info.model_card_template = model_card_template250 if repo_url is not None:251 info.repo_url = repo_url252 if paper_url is not None:253 info.paper_url = paper_url254 if docs_url is not None:255 info.docs_url = docs_url256 if language is not None:257 info.model_card_data.language = language258 if library_name is not None:259 info.model_card_data.library_name = library_name260 if license is not None:261 info.model_card_data.license = license262 if license_name is not None:263 info.model_card_data.license_name = license_name264 if license_link is not None:265 info.model_card_data.license_link = license_link266 if pipeline_tag is not None:267 info.model_card_data.pipeline_tag = pipeline_tag268 if tags is not None:269 normalized_tags = list(tags)270 if info.model_card_data.tags is not None:271 info.model_card_data.tags.extend(normalized_tags)272 else:273 info.model_card_data.tags = normalized_tags274 275 if info.model_card_data.tags is not None:276 info.model_card_data.tags = sorted(set(info.model_card_data.tags))277 278 # Handle encoders/decoders for args279 cls._hub_mixin_coders = coders or {}280 cls._hub_mixin_jsonable_custom_types = tuple(cls._hub_mixin_coders.keys())281 282 # Inspect __init__ signature to handle config283 cls._hub_mixin_init_parameters = dict(inspect.signature(cls.__init__).parameters)284 cls._hub_mixin_jsonable_default_values = {285 param.name: cls._encode_arg(param.default)286 for param in cls._hub_mixin_init_parameters.values()287 if param.default is not inspect.Parameter.empty and cls._is_jsonable(param.default)288 }289 cls._hub_mixin_inject_config = "config" in inspect.signature(cls._from_pretrained).parameters290 291 def __new__(cls: type[T], *args, **kwargs) -> T:292 """Create a new instance of the class and handle config.293 294 3 cases:295 - If `self._hub_mixin_config` is already set, do nothing.296 - If `config` is passed as a dataclass, set it as `self._hub_mixin_config`.297 - Otherwise, build `self._hub_mixin_config` from default values and passed values.298 """299 instance = super().__new__(cls)300 301 # If `config` is already set, return early302 if instance._hub_mixin_config is not None:303 return instance304 305 # Infer passed values306 passed_values = {307 **{308 key: value309 for key, value in zip(310 # [1:] to skip `self` parameter311 list(cls._hub_mixin_init_parameters)[1:],312 args,313 )314 },315 **kwargs,316 }317 318 # If config passed as dataclass => set it and return early319 if is_dataclass(passed_values.get("config")):320 instance._hub_mixin_config = passed_values["config"]321 return instance322 323 # Otherwise, build config from default + passed values324 init_config = {325 # default values326 **cls._hub_mixin_jsonable_default_values,327 # passed values328 **{329 key: cls._encode_arg(value) # Encode custom types as jsonable value330 for key, value in passed_values.items()331 if instance._is_jsonable(value) # Only if jsonable or we have a custom encoder332 },333 }334 passed_config = init_config.pop("config", {})335 336 # Populate `init_config` with provided config337 if isinstance(passed_config, dict):338 init_config.update(passed_config)339 340 # Set `config` attribute and return341 if init_config != {}:342 instance._hub_mixin_config = init_config343 return instance344 345 @classmethod346 def _is_jsonable(cls, value: Any) -> bool:347 """Check if a value is JSON serializable."""348 if is_dataclass(value):349 return True350 if isinstance(value, cls._hub_mixin_jsonable_custom_types):351 return True352 return is_jsonable(value)353 354 @classmethod355 def _encode_arg(cls, arg: Any) -> Any:356 """Encode an argument into a JSON serializable format."""357 if is_dataclass(arg):358 return asdict(arg) # type: ignore[arg-type]359 for type_, (encoder, _) in cls._hub_mixin_coders.items():360 if isinstance(arg, type_):361 if arg is None:362 return None363 return encoder(arg)364 return arg365 366 @classmethod367 def _decode_arg(cls, expected_type: type[ARGS_T], value: Any) -> ARGS_T | None:368 """Decode a JSON serializable value into an argument."""369 if is_simple_optional_type(expected_type):370 if value is None:371 return None372 expected_type = unwrap_simple_optional_type(expected_type) # type: ignore373 # Dataclass => handle it374 if is_dataclass(expected_type):375 return _load_dataclass(expected_type, value) # type: ignore376 # Otherwise => check custom decoders377 for type_, (_, decoder) in cls._hub_mixin_coders.items():378 if inspect.isclass(expected_type) and issubclass(expected_type, type_):379 return decoder(value)380 # Otherwise => don't decode381 return value382 383 def save_pretrained(384 self,385 save_directory: str | Path,386 *,387 config: dict | DataclassInstance | None = None,388 repo_id: str | None = None,389 push_to_hub: bool = False,390 model_card_kwargs: dict[str, Any] | None = None,391 **push_to_hub_kwargs,392 ) -> str | None:393 """394 Save weights in local directory.395 396 Args:397 save_directory (`str` or `Path`):398 Path to directory in which the model weights and configuration will be saved.399 config (`dict` or `DataclassInstance`, *optional*):400 Model configuration specified as a key/value dictionary or a dataclass instance.401 push_to_hub (`bool`, *optional*, defaults to `False`):402 Whether or not to push your model to the Huggingface Hub after saving it.403 repo_id (`str`, *optional*):404 ID of your repository on the Hub. Used only if `push_to_hub=True`. Will default to the folder name if405 not provided.406 model_card_kwargs (`dict[str, Any]`, *optional*):407 Additional arguments passed to the model card template to customize the model card.408 push_to_hub_kwargs:409 Additional key word arguments passed along to the [`~ModelHubMixin.push_to_hub`] method.410 Returns:411 `str` or `None`: url of the commit on the Hub if `push_to_hub=True`, `None` otherwise.412 """413 save_directory = Path(save_directory)414 save_directory.mkdir(parents=True, exist_ok=True)415 416 # Remove config.json if already exists. After `_save_pretrained` we don't want to overwrite config.json417 # as it might have been saved by the custom `_save_pretrained` already. However we do want to overwrite418 # an existing config.json if it was not saved by `_save_pretrained`.419 config_path = save_directory / constants.CONFIG_NAME420 config_path.unlink(missing_ok=True)421 422 # save model weights/files (framework-specific)423 self._save_pretrained(save_directory)424 425 # save config (if provided and if not serialized yet in `_save_pretrained`)426 if config is None:427 config = self._hub_mixin_config428 if config is not None:429 if is_dataclass(config):430 config = asdict(config) # type: ignore[arg-type]431 if not config_path.exists():432 config_str = json.dumps(config, sort_keys=True, indent=2)433 config_path.write_text(config_str)434 435 # save model card436 model_card_path = save_directory / "README.md"437 model_card_kwargs = model_card_kwargs if model_card_kwargs is not None else {}438 if not model_card_path.exists(): # do not overwrite if already exists439 self.generate_model_card(**model_card_kwargs).save(save_directory / "README.md")440 441 # push to the Hub if required442 if push_to_hub:443 kwargs = push_to_hub_kwargs.copy() # soft-copy to avoid mutating input444 if config is not None: # kwarg for `push_to_hub`445 kwargs["config"] = config446 if repo_id is None:447 repo_id = save_directory.name # Defaults to `save_directory` name448 return self.push_to_hub(repo_id=repo_id, model_card_kwargs=model_card_kwargs, **kwargs)449 return None450 451 def _save_pretrained(self, save_directory: Path) -> None:452 """453 Overwrite this method in subclass to define how to save your model.454 Check out our [integration guide](../guides/integrations) for instructions.455 456 Args:457 save_directory (`str` or `Path`):458 Path to directory in which the model weights and configuration will be saved.459 """460 raise NotImplementedError461 462 @classmethod463 @validate_hf_hub_args464 def from_pretrained(465 cls: type[T],466 pretrained_model_name_or_path: str | Path,467 *,468 force_download: bool = False,469 token: str | bool | None = None,470 cache_dir: str | Path | None = None,471 local_files_only: bool = False,472 revision: str | None = None,473 **model_kwargs,474 ) -> T:475 """476 Download a model from the Huggingface Hub and instantiate it.477 478 Args:479 pretrained_model_name_or_path (`str`, `Path`):480 - Either the `model_id` (string) of a model hosted on the Hub, e.g. `bigscience/bloom`.481 - Or a path to a `directory` containing model weights saved using482 [`~transformers.PreTrainedModel.save_pretrained`], e.g., `../path/to/my_model_directory/`.483 revision (`str`, *optional*):484 Revision of the model on the Hub. Can be a branch name, a git tag or any commit id.485 Defaults to the latest commit on `main` branch.486 force_download (`bool`, *optional*, defaults to `False`):487 Whether to force (re-)downloading the model weights and configuration files from the Hub, overriding488 the existing cache.489 token (`str` or `bool`, *optional*):490 The token to use as HTTP bearer authorization for remote files. By default, it will use the token491 cached when running `hf auth login`.492 cache_dir (`str`, `Path`, *optional*):493 Path to the folder where cached files are stored.494 local_files_only (`bool`, *optional*, defaults to `False`):495 If `True`, avoid downloading the file and return the path to the local cached file if it exists.496 model_kwargs (`dict`, *optional*):497 Additional kwargs to pass to the model during initialization.498 """499 model_id = str(pretrained_model_name_or_path)500 config_file: str | None = None501 if os.path.isdir(model_id):502 if constants.CONFIG_NAME in os.listdir(model_id):503 config_file = os.path.join(model_id, constants.CONFIG_NAME)504 else:505 logger.warning(f"{constants.CONFIG_NAME} not found in {Path(model_id).resolve()}")506 else:507 try:508 config_file = hf_hub_download(509 repo_id=model_id,510 filename=constants.CONFIG_NAME,511 revision=revision,512 cache_dir=cache_dir,513 force_download=force_download,514 token=token,515 local_files_only=local_files_only,516 )517 except HfHubHTTPError as e:518 logger.info(f"{constants.CONFIG_NAME} not found on the HuggingFace Hub: {str(e)}")519 520 # Read config521 config = None522 if config_file is not None:523 with open(config_file, encoding="utf-8") as f:524 config = json.load(f)525 526 # Decode custom types in config527 for key, value in config.items():528 if key in cls._hub_mixin_init_parameters:529 expected_type = cls._hub_mixin_init_parameters[key].annotation530 if expected_type is not inspect.Parameter.empty:531 config[key] = cls._decode_arg(expected_type, value)532 533 # Populate model_kwargs from config534 for param in cls._hub_mixin_init_parameters.values():535 if param.name not in model_kwargs and param.name in config:536 model_kwargs[param.name] = config[param.name]537 538 # Check if `config` argument was passed at init539 if "config" in cls._hub_mixin_init_parameters and "config" not in model_kwargs:540 # Decode `config` argument if it was passed541 config_annotation = cls._hub_mixin_init_parameters["config"].annotation542 config = cls._decode_arg(config_annotation, config)543 544 # Forward config to model initialization545 model_kwargs["config"] = config546 547 # Inject config if `**kwargs` are expected548 if is_dataclass(cls):549 for key in cls.__dataclass_fields__:550 if key not in model_kwargs and key in config:551 model_kwargs[key] = config[key]552 elif any(param.kind == inspect.Parameter.VAR_KEYWORD for param in cls._hub_mixin_init_parameters.values()):553 for key, value in config.items(): # type: ignore[union-attr]554 if key not in model_kwargs:555 model_kwargs[key] = value556 557 # Finally, also inject if `_from_pretrained` expects it558 if cls._hub_mixin_inject_config and "config" not in model_kwargs:559 model_kwargs["config"] = config560 561 instance = cls._from_pretrained(562 model_id=str(model_id),563 revision=revision,564 cache_dir=cache_dir,565 force_download=force_download,566 local_files_only=local_files_only,567 token=token,568 **model_kwargs,569 )570 571 # Implicitly set the config as instance attribute if not already set by the class572 # This way `config` will be available when calling `save_pretrained` or `push_to_hub`.573 if config is not None and (getattr(instance, "_hub_mixin_config", None) in (None, {})):574 instance._hub_mixin_config = config575 576 return instance577 578 @classmethod579 def _from_pretrained(580 cls: type[T],581 *,582 model_id: str,583 revision: str | None,584 cache_dir: str | Path | None,585 force_download: bool,586 local_files_only: bool,587 token: str | bool | None,588 **model_kwargs,589 ) -> T:590 """Overwrite this method in subclass to define how to load your model from pretrained.591 592 Use [`hf_hub_download`] or [`snapshot_download`] to download files from the Hub before loading them. Most593 args taken as input can be directly passed to those 2 methods. If needed, you can add more arguments to this594 method using "model_kwargs". For example [`PyTorchModelHubMixin._from_pretrained`] takes as input a `map_location`595 parameter to set on which device the model should be loaded.596 597 Check out our [integration guide](../guides/integrations) for more instructions.598 599 Args:600 model_id (`str`):601 ID of the model to load from the Huggingface Hub (e.g. `bigscience/bloom`).602 revision (`str`, *optional*):603 Revision of the model on the Hub. Can be a branch name, a git tag or any commit id. Defaults to the604 latest commit on `main` branch.605 force_download (`bool`, *optional*, defaults to `False`):606 Whether to force (re-)downloading the model weights and configuration files from the Hub, overriding607 the existing cache.608 token (`str` or `bool`, *optional*):609 The token to use as HTTP bearer authorization for remote files. By default, it will use the token610 cached when running `hf auth login`.611 cache_dir (`str`, `Path`, *optional*):612 Path to the folder where cached files are stored.613 local_files_only (`bool`, *optional*, defaults to `False`):614 If `True`, avoid downloading the file and return the path to the local cached file if it exists.615 model_kwargs:616 Additional keyword arguments passed along to the [`~ModelHubMixin._from_pretrained`] method.617 """618 raise NotImplementedError619 620 @validate_hf_hub_args621 def push_to_hub(622 self,623 repo_id: str,624 *,625 config: dict | DataclassInstance | None = None,626 commit_message: str = "Push model using huggingface_hub.",627 private: bool | None = None,628 token: str | None = None,629 branch: str | None = None,630 create_pr: bool | None = None,631 allow_patterns: list[str] | str | None = None,632 ignore_patterns: list[str] | str | None = None,633 delete_patterns: list[str] | str | None = None,634 model_card_kwargs: dict[str, Any] | None = None,635 ) -> str:636 """637 Upload model checkpoint to the Hub.638 639 Use `allow_patterns` and `ignore_patterns` to precisely filter which files should be pushed to the hub. Use640 `delete_patterns` to delete existing remote files in the same commit. See [`upload_folder`] reference for more641 details.642 643 Args:644 repo_id (`str`):645 ID of the repository to push to (example: `"username/my-model"`).646 config (`dict` or `DataclassInstance`, *optional*):647 Model configuration specified as a key/value dictionary or a dataclass instance.648 commit_message (`str`, *optional*):649 Message to commit while pushing.650 private (`bool`, *optional*):651 Whether the repository created should be private.652 If `None` (default), the repo will be public unless the organization's default is private.653 token (`str`, *optional*):654 The token to use as HTTP bearer authorization for remote files. By default, it will use the token655 cached when running `hf auth login`.656 branch (`str`, *optional*):657 The git branch on which to push the model. This defaults to `"main"`.658 create_pr (`boolean`, *optional*):659 Whether or not to create a Pull Request from `branch` with that commit. Defaults to `False`.660 allow_patterns (`list[str]` or `str`, *optional*):661 If provided, only files matching at least one pattern are pushed.662 ignore_patterns (`list[str]` or `str`, *optional*):663 If provided, files matching any of the patterns are not pushed.664 delete_patterns (`list[str]` or `str`, *optional*):665 If provided, remote files matching any of the patterns will be deleted from the repo.666 model_card_kwargs (`dict[str, Any]`, *optional*):667 Additional arguments passed to the model card template to customize the model card.668 669 Returns:670 The url of the commit of your model in the given repository.671 """672 api = HfApi(token=token)673 repo_id = api.create_repo(repo_id=repo_id, private=private, exist_ok=True).repo_id674 675 # Push the files to the repo in a single commit676 with SoftTemporaryDirectory() as tmp:677 saved_path = Path(tmp) / repo_id678 self.save_pretrained(saved_path, config=config, model_card_kwargs=model_card_kwargs)679 return api.upload_folder(680 repo_id=repo_id,681 repo_type="model",682 folder_path=saved_path,683 commit_message=commit_message,684 revision=branch,685 create_pr=create_pr,686 allow_patterns=allow_patterns,687 ignore_patterns=ignore_patterns,688 delete_patterns=delete_patterns,689 )690 691 def generate_model_card(self, *args, **kwargs) -> ModelCard:692 card = ModelCard.from_template(693 card_data=self._hub_mixin_info.model_card_data,694 template_str=self._hub_mixin_info.model_card_template,695 repo_url=self._hub_mixin_info.repo_url,696 paper_url=self._hub_mixin_info.paper_url,697 docs_url=self._hub_mixin_info.docs_url,698 **kwargs,699 )700 return card701 702 703class PyTorchModelHubMixin(ModelHubMixin):704 """705 Implementation of [`ModelHubMixin`] to provide model Hub upload/download capabilities to PyTorch models. The model706 is set in evaluation mode by default using `model.eval()` (dropout modules are deactivated). To train the model,707 you should first set it back in training mode with `model.train()`.708 709 See [`ModelHubMixin`] for more details on how to use the mixin.710 711 Example:712 713 ```python714 >>> import torch715 >>> import torch.nn as nn716 >>> from huggingface_hub import PyTorchModelHubMixin717 718 >>> class MyModel(719 ... nn.Module,720 ... PyTorchModelHubMixin,721 ... library_name="keras-nlp",722 ... repo_url="https://github.com/keras-team/keras-nlp",723 ... paper_url="https://arxiv.org/abs/2304.12244",724 ... docs_url="https://keras.io/keras_nlp/",725 ... # ^ optional metadata to generate model card726 ... ):727 ... def __init__(self, hidden_size: int = 512, vocab_size: int = 30000, output_size: int = 4):728 ... super().__init__()729 ... self.param = nn.Parameter(torch.rand(hidden_size, vocab_size))730 ... self.linear = nn.Linear(output_size, vocab_size)731 732 ... def forward(self, x):733 ... return self.linear(x + self.param)734 >>> model = MyModel(hidden_size=256)735 736 # Save model weights to local directory737 >>> model.save_pretrained("my-awesome-model")738 739 # Push model weights to the Hub740 >>> model.push_to_hub("my-awesome-model")741 742 # Download and initialize weights from the Hub743 >>> model = MyModel.from_pretrained("username/my-awesome-model")744 >>> model.hidden_size745 256746 ```747 """748 749 def __init_subclass__(cls, *args, tags: list[str] | None = None, **kwargs) -> None:750 tags = tags or []751 tags.append("pytorch_model_hub_mixin")752 kwargs["tags"] = tags753 return super().__init_subclass__(*args, **kwargs)754 755 def _save_pretrained(self, save_directory: Path) -> None:756 """Save weights from a Pytorch model to a local directory."""757 model_to_save = self.module if hasattr(self, "module") else self # type: ignore758 save_model_as_safetensor(model_to_save, str(save_directory / constants.SAFETENSORS_SINGLE_FILE)) # type: ignore [arg-type]759 760 @classmethod761 def _from_pretrained(762 cls,763 *,764 model_id: str,765 revision: str | None,766 cache_dir: str | Path | None,767 force_download: bool,768 local_files_only: bool,769 token: str | bool | None,770 map_location: str = "cpu",771 strict: bool = False,772 **model_kwargs,773 ):774 """Load Pytorch pretrained weights and return the loaded model."""775 model = cls(**model_kwargs)776 if os.path.isdir(model_id):777 print("Loading weights from local directory")778 model_file = os.path.join(model_id, constants.SAFETENSORS_SINGLE_FILE)779 return cls._load_as_safetensor(model, model_file, map_location, strict)780 else:781 try:782 model_file = hf_hub_download(783 repo_id=model_id,784 filename=constants.SAFETENSORS_SINGLE_FILE,785 revision=revision,786 cache_dir=cache_dir,787 force_download=force_download,788 token=token,789 local_files_only=local_files_only,790 )791 return cls._load_as_safetensor(model, model_file, map_location, strict)792 except EntryNotFoundError:793 model_file = hf_hub_download(794 repo_id=model_id,795 filename=constants.PYTORCH_WEIGHTS_NAME,796 revision=revision,797 cache_dir=cache_dir,798 force_download=force_download,799 token=token,800 local_files_only=local_files_only,801 )802 return cls._load_as_pickle(model, model_file, map_location, strict)803 804 @classmethod805 def _load_as_pickle(cls, model: T, model_file: str, map_location: str, strict: bool) -> T:806 state_dict = torch.load(model_file, map_location=torch.device(map_location), weights_only=True)807 model.load_state_dict(state_dict, strict=strict) # type: ignore808 model.eval() # type: ignore809 return model810 811 @classmethod812 def _load_as_safetensor(cls, model: T, model_file: str, map_location: str, strict: bool) -> T:813 if packaging.version.parse(safetensors.__version__) < packaging.version.parse("0.4.3"): # type: ignore [attr-defined]814 load_model_as_safetensor(model, model_file, strict=strict) # type: ignore [arg-type]815 if map_location != "cpu":816 logger.warning(817 "Loading model weights on other devices than 'cpu' is not supported natively in your version of safetensors."818 " This means that the model is loaded on 'cpu' first and then copied to the device."819 " This leads to a slower loading time."820 " Please update safetensors to version 0.4.3 or above for improved performance."821 )822 model.to(map_location) # type: ignore [attr-defined]823 else:824 safetensors.torch.load_model(model, model_file, strict=strict, device=map_location) # type: ignore [arg-type]825 model.eval() # type: ignore826 return model827 828 829def _load_dataclass(datacls: type[DataclassInstance], data: dict) -> DataclassInstance:830 """Load a dataclass instance from a dictionary.831 832 Fields not expected by the dataclass are ignored.833 """834 return datacls(**{k: v for k, v in data.items() if k in datacls.__dataclass_fields__})835 