Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
utils.py182 linesDownload Raw Back to models
1import torch, os2from safetensors import safe_open3from contextlib import contextmanager4import hashlib5 6@contextmanager7def init_weights_on_device(device = torch.device("meta"), include_buffers :bool = False):8    9    old_register_parameter = torch.nn.Module.register_parameter10    if include_buffers:11        old_register_buffer = torch.nn.Module.register_buffer12    13    def register_empty_parameter(module, name, param):14        old_register_parameter(module, name, param)15        if param is not None:16            param_cls = type(module._parameters[name])17            kwargs = module._parameters[name].__dict__18            kwargs["requires_grad"] = param.requires_grad19            module._parameters[name] = param_cls(module._parameters[name].to(device), **kwargs)20 21    def register_empty_buffer(module, name, buffer, persistent=True):22        old_register_buffer(module, name, buffer, persistent=persistent)23        if buffer is not None:24            module._buffers[name] = module._buffers[name].to(device)25            26    def patch_tensor_constructor(fn):27        def wrapper(*args, **kwargs):28            kwargs["device"] = device29            return fn(*args, **kwargs)30 31        return wrapper32    33    if include_buffers:34        tensor_constructors_to_patch = {35            torch_function_name: getattr(torch, torch_function_name)36            for torch_function_name in ["empty", "zeros", "ones", "full"]37        }38    else:39        tensor_constructors_to_patch = {}40    41    try:42        torch.nn.Module.register_parameter = register_empty_parameter43        if include_buffers:44            torch.nn.Module.register_buffer = register_empty_buffer45        for torch_function_name in tensor_constructors_to_patch.keys():46            setattr(torch, torch_function_name, patch_tensor_constructor(getattr(torch, torch_function_name)))47        yield48    finally:49        torch.nn.Module.register_parameter = old_register_parameter50        if include_buffers:51            torch.nn.Module.register_buffer = old_register_buffer52        for torch_function_name, old_torch_function in tensor_constructors_to_patch.items():53            setattr(torch, torch_function_name, old_torch_function)54 55def load_state_dict_from_folder(file_path, torch_dtype=None):56    state_dict = {}57    for file_name in os.listdir(file_path):58        if "." in file_name and file_name.split(".")[-1] in [59            "safetensors", "bin", "ckpt", "pth", "pt"60        ]:61            state_dict.update(load_state_dict(os.path.join(file_path, file_name), torch_dtype=torch_dtype))62    return state_dict63 64 65def load_state_dict(file_path, torch_dtype=None, device="cpu"):66    if file_path.endswith(".safetensors"):67        return load_state_dict_from_safetensors(file_path, torch_dtype=torch_dtype, device=device)68    else:69        return load_state_dict_from_bin(file_path, torch_dtype=torch_dtype, device=device)70 71 72def load_state_dict_from_safetensors(file_path, torch_dtype=None, device="cpu"):73    state_dict = {}74    with safe_open(file_path, framework="pt", device=str(device)) as f:75        for k in f.keys():76            state_dict[k] = f.get_tensor(k)77            if torch_dtype is not None:78                state_dict[k] = state_dict[k].to(torch_dtype)79    return state_dict80 81 82def load_state_dict_from_bin(file_path, torch_dtype=None, device="cpu"):83    state_dict = torch.load(file_path, map_location=device, weights_only=True)84    if torch_dtype is not None:85        for i in state_dict:86            if isinstance(state_dict[i], torch.Tensor):87                state_dict[i] = state_dict[i].to(torch_dtype)88    return state_dict89 90 91def search_for_embeddings(state_dict):92    embeddings = []93    for k in state_dict:94        if isinstance(state_dict[k], torch.Tensor):95            embeddings.append(state_dict[k])96        elif isinstance(state_dict[k], dict):97            embeddings += search_for_embeddings(state_dict[k])98    return embeddings99 100 101def search_parameter(param, state_dict):102    for name, param_ in state_dict.items():103        if param.numel() == param_.numel():104            if param.shape == param_.shape:105                if torch.dist(param, param_) < 1e-3:106                    return name107            else:108                if torch.dist(param.flatten(), param_.flatten()) < 1e-3:109                    return name110    return None111 112 113def build_rename_dict(source_state_dict, target_state_dict, split_qkv=False):114    matched_keys = set()115    with torch.no_grad():116        for name in source_state_dict:117            rename = search_parameter(source_state_dict[name], target_state_dict)118            if rename is not None:119                print(f'"{name}": "{rename}",')120                matched_keys.add(rename)121            elif split_qkv and len(source_state_dict[name].shape)>=1 and source_state_dict[name].shape[0]%3==0:122                length = source_state_dict[name].shape[0] // 3123                rename = []124                for i in range(3):125                    rename.append(search_parameter(source_state_dict[name][i*length: i*length+length], target_state_dict))126                if None not in rename:127                    print(f'"{name}": {rename},')128                    for rename_ in rename:129                        matched_keys.add(rename_)130    for name in target_state_dict:131        if name not in matched_keys:132            print("Cannot find", name, target_state_dict[name].shape)133 134 135def search_for_files(folder, extensions):136    files = []137    if os.path.isdir(folder):138        for file in sorted(os.listdir(folder)):139            files += search_for_files(os.path.join(folder, file), extensions)140    elif os.path.isfile(folder):141        for extension in extensions:142            if folder.endswith(extension):143                files.append(folder)144                break145    return files146 147 148def convert_state_dict_keys_to_single_str(state_dict, with_shape=True):149    keys = []150    for key, value in state_dict.items():151        if isinstance(key, str):152            if isinstance(value, torch.Tensor):153                if with_shape:154                    shape = "_".join(map(str, list(value.shape)))155                    keys.append(key + ":" + shape)156                keys.append(key)157            elif isinstance(value, dict):158                keys.append(key + "|" + convert_state_dict_keys_to_single_str(value, with_shape=with_shape))159    keys.sort()160    keys_str = ",".join(keys)161    return keys_str162 163 164def split_state_dict_with_prefix(state_dict):165    keys = sorted([key for key in state_dict if isinstance(key, str)])166    prefix_dict = {}167    for key in  keys:168        prefix = key if "." not in key else key.split(".")[0]169        if prefix not in prefix_dict:170            prefix_dict[prefix] = []171        prefix_dict[prefix].append(key)172    state_dicts = []173    for prefix, keys in prefix_dict.items():174        sub_state_dict = {key: state_dict[key] for key in keys}175        state_dicts.append(sub_state_dict)176    return state_dicts177 178 179def hash_state_dict_keys(state_dict, with_shape=True):180    keys_str = convert_state_dict_keys_to_single_str(state_dict, with_shape=with_shape)181    keys_str = keys_str.encode(encoding="UTF-8")182    return hashlib.md5(keys_str).hexdigest()