Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_base.py209 linesDownload Raw Back to serialization
1# Copyright 2024 The HuggingFace Team. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7#     http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14"""Contains helpers to split tensors into shards."""15 16from collections.abc import Callable17from dataclasses import dataclass, field18from typing import Any, TypeVar19 20from .. import logging21 22 23TensorT = TypeVar("TensorT")24TensorSizeFn_T = Callable[[TensorT], int]25StorageIDFn_T = Callable[[TensorT], Any | None]26 27MAX_SHARD_SIZE = "5GB"28SIZE_UNITS = {29    "TB": 10**12,30    "GB": 10**9,31    "MB": 10**6,32    "KB": 10**3,33}34 35 36logger = logging.get_logger(__file__)37 38 39@dataclass40class StateDictSplit:41    is_sharded: bool = field(init=False)42    metadata: dict[str, Any]43    filename_to_tensors: dict[str, list[str]]44    tensor_to_filename: dict[str, str]45 46    def __post_init__(self):47        self.is_sharded = len(self.filename_to_tensors) > 148 49 50def split_state_dict_into_shards_factory(51    state_dict: dict[str, TensorT],52    *,53    get_storage_size: TensorSizeFn_T,54    filename_pattern: str,55    get_storage_id: StorageIDFn_T = lambda tensor: None,56    max_shard_size: int | str = MAX_SHARD_SIZE,57) -> StateDictSplit:58    """59    Split a model state dictionary in shards so that each shard is smaller than a given size.60 61    The shards are determined by iterating through the `state_dict` in the order of its keys. There is no optimization62    made to make each shard as close as possible to the maximum size passed. For example, if the limit is 10GB and we63    have tensors of sizes [6GB, 6GB, 2GB, 6GB, 2GB, 2GB] they will get sharded as [6GB], [6+2GB], [6+2+2GB] and not64    [6+2+2GB], [6+2GB], [6GB].65 66    > [!WARNING]67    > If one of the model's tensor is bigger than `max_shard_size`, it will end up in its own shard which will have a68    > size greater than `max_shard_size`.69 70    Args:71        state_dict (`dict[str, Tensor]`):72            The state dictionary to save.73        get_storage_size (`Callable[[Tensor], int]`):74            A function that returns the size of a tensor when saved on disk in bytes.75        get_storage_id (`Callable[[Tensor], Optional[Any]]`, *optional*):76            A function that returns a unique identifier to a tensor storage. Multiple different tensors can share the77            same underlying storage. This identifier is guaranteed to be unique and constant for this tensor's storage78            during its lifetime. Two tensor storages with non-overlapping lifetimes may have the same id.79        filename_pattern (`str`, *optional*):80            The pattern to generate the files names in which the model will be saved. Pattern must be a string that81            can be formatted with `filename_pattern.format(suffix=...)` and must contain the keyword `suffix`82        max_shard_size (`int` or `str`, *optional*):83            The maximum size of each shard, in bytes. Defaults to 5GB.84 85    Returns:86        [`StateDictSplit`]: A `StateDictSplit` object containing the shards and the index to retrieve them.87    """88    storage_id_to_tensors: dict[Any, list[str]] = {}89 90    shard_list: list[dict[str, TensorT]] = []91    current_shard: dict[str, TensorT] = {}92    current_shard_size = 093    total_size = 094 95    if isinstance(max_shard_size, str):96        max_shard_size = parse_size_to_int(max_shard_size)97 98    for key, tensor in state_dict.items():99        # when bnb serialization is used the weights in the state dict can be strings100        # check: https://github.com/huggingface/transformers/pull/24416 for more details101        if isinstance(tensor, str):102            logger.info("Skipping tensor %s as it is a string (bnb serialization)", key)103            continue104 105        # If a `tensor` shares the same underlying storage as another tensor, we put `tensor` in the same `block`106        storage_id = get_storage_id(tensor)  # type: ignore[invalid-argument-type]107        if storage_id is not None:108            if storage_id in storage_id_to_tensors:109                # We skip this tensor for now and will reassign to correct shard later110                storage_id_to_tensors[storage_id].append(key)111                continue112            else:113                # This is the first tensor with this storage_id, we create a new entry114                # in the storage_id_to_tensors dict => we will assign the shard id later115                storage_id_to_tensors[storage_id] = [key]116 117        # Compute tensor size118        tensor_size = get_storage_size(tensor)  # type: ignore[invalid-argument-type]119 120        # If this tensor is bigger than the maximal size, we put it in its own shard121        if tensor_size > max_shard_size:122            total_size += tensor_size123            shard_list.append({key: tensor})124            continue125 126        # If this tensor is going to tip up over the maximal size, we split.127        # Current shard already has some tensors, we add it to the list of shards and create a new one.128        if current_shard_size + tensor_size > max_shard_size:129            shard_list.append(current_shard)130            current_shard = {}131            current_shard_size = 0132 133        # Add the tensor to the current shard134        current_shard[key] = tensor135        current_shard_size += tensor_size136        total_size += tensor_size137 138    # Add the last shard139    if len(current_shard) > 0:140        shard_list.append(current_shard)141    nb_shards = len(shard_list)142 143    # Loop over the tensors that share the same storage and assign them together144    for storage_id, keys in storage_id_to_tensors.items():145        # Let's try to find the shard where the first tensor of this storage is and put all tensors in the same shard146        for shard in shard_list:147            if keys[0] in shard:148                for key in keys:149                    shard[key] = state_dict[key]150                break151 152    # If we only have one shard, we return it => no need to build the index153    if nb_shards == 1:154        filename = filename_pattern.format(suffix="")155        return StateDictSplit(156            metadata={"total_size": total_size},157            filename_to_tensors={filename: list(state_dict.keys())},158            tensor_to_filename={key: filename for key in state_dict.keys()},159        )160 161    # Now that each tensor is assigned to a shard, let's assign a filename to each shard162    tensor_name_to_filename = {}163    filename_to_tensors = {}164    for idx, shard in enumerate(shard_list):165        filename = filename_pattern.format(suffix=f"-{idx + 1:05d}-of-{nb_shards:05d}")166        for key in shard:167            tensor_name_to_filename[key] = filename168        filename_to_tensors[filename] = list(shard.keys())169 170    # Build the index and return171    return StateDictSplit(172        metadata={"total_size": total_size},173        filename_to_tensors=filename_to_tensors,174        tensor_to_filename=tensor_name_to_filename,175    )176 177 178def parse_size_to_int(size_as_str: str) -> int:179    """180    Parse a size expressed as a string with digits and unit (like `"5MB"`) to an integer (in bytes).181 182    Supported units are "TB", "GB", "MB", "KB".183 184    Args:185        size_as_str (`str`): The size to convert. Will be directly returned if an `int`.186 187    Example:188 189    ```py190    >>> parse_size_to_int("5MB")191    5000000192    ```193    """194    size_as_str = size_as_str.strip()195 196    # Parse unit197    unit = size_as_str[-2:].upper()198    if unit not in SIZE_UNITS:199        raise ValueError(f"Unit '{unit}' not supported. Supported units are TB, GB, MB, KB. Got '{size_as_str}'.")200    multiplier = SIZE_UNITS[unit]201 202    # Parse value203    try:204        value = float(size_as_str[:-2].strip())205    except ValueError as e:206        raise ValueError(f"Could not parse the size value from '{size_as_str}': {e}") from e207 208    return int(value * multiplier)209 
codekingpro/portable-devtools · Team Ai