Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
hub_mixin.py835 linesDownload Raw Back to huggingface_hub
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 
codekingpro/portable-devtools · Team Ai