hugging-apps/echo-memory
0
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()