| 158 | |
| 159 | @dataclass |
| 160 | class ModelConfig: |
| 161 | path: Union[str, list[str]] = None |
| 162 | model_id: str = None |
| 163 | origin_file_pattern: Union[str, list[str]] = None |
| 164 | download_resource: str = "ModelScope" |
| 165 | offload_device: Optional[Union[str, torch.device]] = None |
| 166 | offload_dtype: Optional[torch.dtype] = None |
| 167 | local_model_path: str = None |
| 168 | skip_download: bool = False |
| 169 | |
| 170 | def download_if_necessary(self, use_usp=False): |
| 171 | if self.path is None: |
| 172 | # Check model_id and origin_file_pattern |
| 173 | if self.model_id is None: |
| 174 | raise ValueError(f"""No valid model files. Please use `ModelConfig(path="xxx")` or `ModelConfig(model_id="xxx/yyy", origin_file_pattern="zzz")`.""") |
| 175 | |
| 176 | # Skip if not in rank 0 |
| 177 | if use_usp: |
| 178 | import torch.distributed as dist |
| 179 | skip_download = self.skip_download or dist.get_rank() != 0 |
| 180 | else: |
| 181 | skip_download = self.skip_download |
| 182 | |
| 183 | # Check whether the origin path is a folder |
| 184 | if self.origin_file_pattern is None or self.origin_file_pattern == "": |
| 185 | self.origin_file_pattern = "" |
| 186 | allow_file_pattern = None |
| 187 | is_folder = True |
| 188 | elif isinstance(self.origin_file_pattern, str) and self.origin_file_pattern.endswith("/"): |
| 189 | allow_file_pattern = self.origin_file_pattern + "*" |
| 190 | is_folder = True |
| 191 | else: |
| 192 | allow_file_pattern = self.origin_file_pattern |
| 193 | is_folder = False |
| 194 | |
| 195 | # Download |
| 196 | if self.local_model_path is None: |
| 197 | self.local_model_path = "./models" |
| 198 | if not skip_download: |
| 199 | downloaded_files = glob.glob(self.origin_file_pattern, root_dir=os.path.join(self.local_model_path, self.model_id)) |
| 200 | if self.download_resource.lower() == "modelscope": |
| 201 | snapshot_download( |
| 202 | self.model_id, |
| 203 | local_dir=os.path.join(self.local_model_path, self.model_id), |
| 204 | allow_file_pattern=allow_file_pattern, |
| 205 | ignore_file_pattern=downloaded_files, |
| 206 | local_files_only=False |
| 207 | ) |
| 208 | elif self.download_resource.lower() == "huggingface": |
| 209 | hf_snapshot_download( |
| 210 | self.model_id, |
| 211 | local_dir=os.path.join(self.local_model_path, self.model_id), |
| 212 | allow_patterns=allow_file_pattern, |
| 213 | ignore_patterns=downloaded_files, |
| 214 | local_files_only=False |
| 215 | ) |
| 216 | else: |
| 217 | raise ValueError("`download_resource` should be `modelscope` or `huggingface`.") |
no outgoing calls
no test coverage detected