codekingpro/portable-devtools
114k
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 