From cdcca50fb44a470bf1c9f21f5fddbde5a17a0558 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 14 Oct 2024 13:03:44 -0400 Subject: [PATCH] add model analyzer Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 12 +++++-- installer.py | 2 ++ javascript/sdnext.css | 1 + modules/model_flux.py | 21 +++++++---- modules/model_te.py | 37 ++++++++++++------- modules/modelstats.py | 82 +++++++++++++++++++++++++++++++++++++++++++ modules/sd_models.py | 4 +++ modules/ui_models.py | 22 ++++++++++++ 8 files changed, 158 insertions(+), 23 deletions(-) create mode 100644 modules/modelstats.py diff --git a/CHANGELOG.md b/CHANGELOG.md index f709bba5f..42d4333a8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,13 +1,15 @@ # Change Log for SD.Next -## Update for 2024-10-12 +## Update for 2024-10-13 -### Highlights for 2024-10-12 +### Highlights for 2024-10-13 - **Reprocess**: New workflow options that allow you to generate at lower quality and then reprocess at higher quality for select images only or generate without hires/refine and then reprocess with hires/refine and you can pick any previous latent from auto-captured history! - **Detailer** Fully built-in detailer workflow without with support for all standard models +- Built-in **model analyzer** + See all details of your currently loaded model, including components, parameter count, layer count, etc. - New fine-tuned [CLiP-ViT-L]((https://huggingface.co/zer0int/CLIP-GmP-ViT-L-14)) 1st stage **text-encoders** used by SD15, SDXL, Flux.1, etc. brings additional details to your images - Integration with [Ctrl+X](https://github.com/genforce/ctrl-x) which allows for control of **structure and appearance** without the need for extra models, [APG: Adaptive Projected Guidance](https://arxiv.org/pdf/2410.02416) for optimal **guidance** control, @@ -19,7 +21,7 @@ And other goodies like multiple *XYZ grid* improvements, additional *Flux ControlNets*, additional *Interrogate models*, better *LoRA tags* support, and more... -### Details for 2024-10-12 +### Details for 2024-10-13 - **reprocess** - new top-level button: reprocess latent from your history of generated image(s) @@ -38,6 +40,10 @@ And other goodies like multiple *XYZ grid* improvements, additional *Flux Contro memory usage is ~130kb of ram for 1mp image - *note* list of latents in history is not auto-refreshed, use refresh button +- **model analyzer** + - see all details of your currently loaded model, including components, parameter count, layer count, etc. + - in models -> current -> analyze + - **text encoder**: - allow loading different custom text encoders: *clip-vit-l, clip-vit-g, t5* will automatically find appropriate encoder in the loaded model and replace it with loaded text encoder diff --git a/installer.py b/installer.py index 3e1e8843b..159aa2e60 100644 --- a/installer.py +++ b/installer.py @@ -166,6 +166,8 @@ def custom_excepthook(exc_type, exc_value, exc_traceback): def print_dict(d): + if d is None: + return '' return ' '.join([f'{k}={v}' for k, v in d.items()]) diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 879999549..0f3765694 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -17,6 +17,7 @@ textarea { overflow-y: auto !important; } span { font-size: var(--text-md) !important; } button { font-size: var(--text-lg) !important; } input[type='color'] { width: 64px; height: 32px; } +td > div > span { overflow-y: auto; max-height: 3em; overflow-x: hidden; } /* gradio elements */ .block .padded:not(.gradio-accordion) { padding: 4px 0 0 0 !important; margin-right: 0; min-width: 90px !important; } diff --git a/modules/model_flux.py b/modules/model_flux.py index 5987d686a..a21d14a0a 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -173,11 +173,12 @@ def load_transformer(file_path): # triggered by opts.sd_unet change def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_checkpoint change quant = model_quant.get_quant(checkpoint_info.path) repo_id = sd_models.path_to_repo(checkpoint_info.name) - shared.log.debug(f'Load model: type=FLUX model="{checkpoint_info.name}" repo="{repo_id}" unet="{shared.opts.sd_unet}" t5="{shared.opts.sd_text_encoder}" vae="{shared.opts.sd_vae}" quant={quant} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}') + shared.log.debug(f'Load model: type=FLUX model="{checkpoint_info.name}" repo="{repo_id}" unet="{shared.opts.sd_unet}" te="{shared.opts.sd_text_encoder}" vae="{shared.opts.sd_vae}" quant={quant} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}') debug(f'Load model: type=FLUX config={diffusers_load_config}') modelloader.hf_login() transformer = None + text_encoder_1 = None text_encoder_2 = None vae = None @@ -199,13 +200,12 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch errors.display(e, 'FLUX UNet:') if shared.opts.sd_text_encoder != 'None': try: - debug(f'Load model: type=FLUX t5="{shared.opts.sd_text_encoder}"') - from modules.model_te import load_t5 - _text_encoder_2 = load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir) - if _text_encoder_2 is not None: - text_encoder_2 = _text_encoder_2 + debug(f'Load model: type=FLUX te="{shared.opts.sd_text_encoder}"') + from modules.model_te import load_t5, load_vit_l + if 'vit-l' in shared.opts.sd_text_encoder.lower(): + text_encoder_1 = load_vit_l() else: - shared.opts.sd_text_encoder = 'None' + text_encoder_2 = load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir) except Exception as e: shared.log.error(f"Load model: type=FLUX Failed to load T5: {e}") shared.opts.sd_text_encoder = 'None' @@ -260,6 +260,9 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch if transformer is not None: components['transformer'] = transformer sd_unet.loaded_unet = shared.opts.sd_unet + if text_encoder_1 is not None: + components['text_encoder'] = text_encoder_1 + model_te.loaded_te = shared.opts.sd_text_encoder if text_encoder_2 is not None: components['text_encoder_2'] = text_encoder_2 model_te.loaded_te = shared.opts.sd_text_encoder @@ -268,6 +271,10 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch shared.log.debug(f'Load model: type=FLUX preloaded={list(components)}') if repo_id == 'sayakpaul/flux.1-dev-nf4': repo_id = 'black-forest-labs/FLUX.1-dev' # workaround since sayakpaul model is missing model_index.json + for c in components: + if components[c].dtype == torch.float32 and devices.dtype != torch.float32: + shared.log.warning(f'Load model: type=FLUX component={c} dtype={components[c].dtype} cast dtype={devices.dtype}') + components[c] = components[c].to(dtype=devices.dtype) pipe = diffusers.FluxPipeline.from_pretrained(repo_id, cache_dir=shared.opts.diffusers_dir, **components, **diffusers_load_config) try: diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["flux"] = diffusers.FluxPipeline diff --git a/modules/model_te.py b/modules/model_te.py index 40d544e1d..606cdb86d 100644 --- a/modules/model_te.py +++ b/modules/model_te.py @@ -127,25 +127,40 @@ def set_t5(pipe, module, t5=None, cache_dir=None): return pipe -def set_clip(pipe): +def load_vit_l(): global loaded_te # pylint: disable=global-statement + config = transformers.PretrainedConfig.from_json_file('configs/sdxl/text_encoder/config.json') + state_dict = load_file(os.path.join(shared.opts.te_dir, f'{shared.opts.sd_text_encoder}.safetensors')) + te = transformers.CLIPTextModel.from_pretrained(pretrained_model_name_or_path=None, state_dict=state_dict, config=config) + loaded_te = shared.opts.sd_text_encoder + te = te.to(dtype=devices.dtype) + return te + + +def load_vit_g(): + global loaded_te # pylint: disable=global-statement + config = transformers.PretrainedConfig.from_json_file('configs/sdxl/text_encoder_2/config.json') + state_dict = load_file(os.path.join(shared.opts.te_dir, f'{shared.opts.sd_text_encoder}.safetensors')) + te = transformers.CLIPTextModelWithProjection.from_pretrained(pretrained_model_name_or_path=None, state_dict=state_dict, config=config) + loaded_te = shared.opts.sd_text_encoder + te = te.to(dtype=devices.dtype) + return te + + +def set_clip(pipe): if loaded_te == shared.opts.sd_text_encoder: return from modules.sd_models import move_model if 'vit-l' in shared.opts.sd_text_encoder.lower() and hasattr(shared.sd_model, 'text_encoder') and shared.sd_model.text_encoder.__class__.__name__ == 'CLIPTextModel': try: - config = transformers.PretrainedConfig.from_json_file('configs/sdxl/text_encoder/config.json') - state_dict = load_file(os.path.join(shared.opts.te_dir, f'{shared.opts.sd_text_encoder}.safetensors')) - te = transformers.CLIPTextModel.from_pretrained(pretrained_model_name_or_path=None, state_dict=state_dict, config=config) + te = load_vit_l() except Exception as e: shared.log.error(f'Load module: type="text_encoder" class="ViT-L" file="{shared.opts.sd_text_encoder}" {e}') if debug: errors.display(e, 'TE:') - state_dict = None te = None if te is not None: - loaded_te = shared.opts.sd_text_encoder - pipe.text_encoder = te.to(dtype=devices.dtype) + pipe.text_encoder = te shared.log.info(f'Load module: type="text_encoder" class="ViT-L" file="{shared.opts.sd_text_encoder}"') import modules.prompt_parser_diffusers modules.prompt_parser_diffusers.cache.clear() @@ -153,18 +168,14 @@ def set_clip(pipe): devices.torch_gc() if 'vit-g' in shared.opts.sd_text_encoder.lower() and hasattr(shared.sd_model, 'text_encoder_2') and shared.sd_model.text_encoder_2.__class__.__name__ == 'CLIPTextModelWithProjection': try: - config = transformers.PretrainedConfig.from_json_file('configs/sdxl/text_encoder_2/config.json') - state_dict = load_file(os.path.join(shared.opts.te_dir, f'{shared.opts.sd_text_encoder}.safetensors')) - te = transformers.CLIPTextModelWithProjection.from_pretrained(pretrained_model_name_or_path=None, state_dict=state_dict, config=config) + te = load_vit_g() except Exception as e: shared.log.error(f'Load module: type module="text_encoder_2" class="ViT-G" file="{shared.opts.sd_text_encoder}" {e}') if debug: errors.display(e, 'TE:') - state_dict = None te = None if te is not None: - loaded_te = shared.opts.sd_text_encoder - pipe.text_encoder_2 = te.to(dtype=devices.dtype) + pipe.text_encoder_2 = te shared.log.info(f'Load module: type="text_encoder_2" class="ViT-G" file="{shared.opts.sd_text_encoder}"') import modules.prompt_parser_diffusers modules.prompt_parser_diffusers.cache.clear() diff --git a/modules/modelstats.py b/modules/modelstats.py new file mode 100644 index 000000000..b0408e10b --- /dev/null +++ b/modules/modelstats.py @@ -0,0 +1,82 @@ +import os +from datetime import datetime +import torch +from modules import shared, sd_models + + +class Module(): + name: str = '' + cls: str = None + device: str = None + dtype: str = None + params: int = 0 + modules: int = 0 + config: dict = None + + def __init__(self, name, module): + self.name = name + # self.type = type(module) + self.cls = module.__class__.__name__ + if hasattr(module, 'config'): + self.config = module.config + if isinstance(module, torch.nn.Module): + self.device = module.device + self.dtype = module.dtype + self.params = sum(p.numel() for p in module.parameters(recurse=True)) + self.modules = len(list(module.modules())) + + def __repr__(self): + s = f'name="{self.name}" cls={self.cls} config={self.config is not None}' + if self.device or self.dtype: + s += f' device={self.device} dtype={self.dtype}' + if self.params or self.modules: + s += f' params={self.params} modules={self.modules}' + return s + + +class Model(): + name: str = '' + fn: str = '' + type: str = '' + cls: str = '' + hash: str = '' + meta: dict = {} + size: int = 0 + mtime: datetime = None + info: sd_models.CheckpointInfo = None + modules: list[Module] = [] + + def __init__(self, name): + self.name = name + if not shared.sd_loaded: + return + self.cls = shared.sd_model.__class__.__name__ + self.type = shared.sd_model_type + self.info = sd_models.get_closet_checkpoint_match(name) + if self.info is not None: + self.name = self.info.name or self.name + self.hash = self.info.shorthash or '' + self.meta = self.info.metadata or {} + if os.path.exists(self.info.filename): + stat = os.stat(self.info.filename) + self.mtime = datetime.fromtimestamp(stat.st_mtime).replace(microsecond=0) + if os.path.isfile(self.info.filename): + self.size = round(stat.st_size) + + def __repr__(self): + return f'model="{self.name}" type={self.type} class={self.cls} size={self.size} mtime="{self.mtime}" modules={self.modules}' + + +def analyze(): + model = Model(shared.opts.sd_model_checkpoint) + if model.cls == '': + return model + if not hasattr(shared.sd_model, '_internal_dict'): + return model + model.modules.clear() + for k in shared.sd_model._internal_dict.keys(): # pylint: disable=protected-access + component = getattr(shared.sd_model, k, None) + module = Module(k, component) + model.modules.append(module) + shared.log.debug(f'Analyzed: {model}') + return model diff --git a/modules/sd_models.py b/modules/sd_models.py index 2a5478be9..155a46a19 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1351,6 +1351,10 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No devices.torch_gc(force=True) if sd_model is not None: script_callbacks.model_loaded_callback(sd_model) + + from modules import modelstats + modelstats.analyze() + shared.log.info(f"Load {op}: time={timer.summary()} native={get_native(sd_model)} memory={memory_stats()}") diff --git a/modules/ui_models.py b/modules/ui_models.py index 64cefc954..fed6441e7 100644 --- a/modules/ui_models.py +++ b/modules/ui_models.py @@ -34,6 +34,28 @@ def create_ui(): def gr_show(visible=True): return {"visible": visible, "__type__": "update"} + with gr.Tab(label="Current"): + def analyze(): + from modules import modelstats + model = modelstats.analyze() + desc = f"Model: {model.name}
Type: {model.type}
Class: {model.cls}
Size: {model.size} bytes
Modified: {model.mtime}
" + meta = model.meta + components = [(m.name, m.cls, m.device, m.dtype, m.params, m.modules, str(m.config)) for m in model.modules] + return [desc, components, meta] + + with gr.Row(): + model_analyze = gr.Button(value="Analyze", variant='primary') + with gr.Row(): + model_desc = gr.HTML(value="", elem_id="model_desc") + with gr.Row(): + module_headers = ['Module', 'Class', 'Device', 'DType', 'Params', 'Modules', 'Config'] + model_types = ['str', 'str', 'str', 'str', 'number', 'number', 'str'] + model_modules = gr.DataFrame(value=None, label=None, show_label=False, interactive=False, wrap=True, headers=module_headers, datatype=model_types, type='array') + with gr.Row(): + model_meta = gr.JSON(label="Metadata", value={}, elem_id="model_meta") + + model_analyze.click(fn=analyze, inputs=[], outputs=[model_desc, model_modules, model_meta]) + with gr.Tab(label="Convert"): with gr.Row(): model_name = gr.Dropdown(sd_models.checkpoint_tiles(), label="Original model")