MCPcopy Create free account
hub / github.com/modelscope/DiffSynth-Studio / ModelConfig

Class ModelConfig

diffsynth/utils/__init__.py:160–230  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

158
159@dataclass
160class 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`.")

Callers 15

from_pretrainedMethod · 0.85
from_pretrainedMethod · 0.85
from_pretrainedMethod · 0.85
parse_model_configsMethod · 0.85
load_modelFunction · 0.85
Step1X-Edit.pyFile · 0.85
FLUX.1-Krea-dev.pyFile · 0.85
FLUX.1-dev.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected