diff --git a/extensions-builtin/Lora/lora_extract.py b/extensions-builtin/Lora/lora_extract.py index 7fa8cab6d..761bf614f 100644 --- a/extensions-builtin/Lora/lora_extract.py +++ b/extensions-builtin/Lora/lora_extract.py @@ -7,6 +7,7 @@ from safetensors.torch import save_file import gradio as gr from modules import shared, devices from modules.ui_common import create_refresh_button +from modules.call_queue import wrap_gradio_gpu_call class SVDHandler: @@ -92,6 +93,7 @@ def make_meta(fn, maxrank, rank_ratio): "model_spec.implementation": "https://github.com/vladmandic/automatic", "model_spec.date": datetime.datetime.now().astimezone().replace(microsecond=0).isoformat(), "model_spec.base_model": shared.opts.sd_model_checkpoint, + "model_spec.dtype": str(devices.dtype), "model_spec.base_lora": json.dumps(loaded_lora()), "model_spec.config": f"maxrank={maxrank} rank_ratio={rank_ratio}", } @@ -128,71 +130,86 @@ def make_lora(fn, maxrank, auto_rank, rank_ratio, modules, overwrite): t0 = time.time() maxrank = int(maxrank) rank_ratio = 1 if not auto_rank else rank_ratio + shared.log.debug(f'LoRA extract: modules={modules} maxrank={maxrank} auto={auto_rank} ratio={rank_ratio} fn="{fn}"') shared.state.begin('LoRA extract') - shared.log.debug(f'LoRA extract: modules={modules} maxrank={maxrank} auto={auto_rank} ratio={rank_ratio} fn="{fn}"') - if 'te' in modules and getattr(shared.sd_model, 'text_encoder', None) is not None: - yield "LoRA extract: extracting TE-1" - for name, module in shared.sd_model.text_encoder.named_modules(): - weights_backup = getattr(module, "network_weights_backup", None) - if weights_backup is None or getattr(module, "network_current_names", None) is None: - continue - prefix = "lora_te1_" if hasattr(shared.sd_model, 'text_encoder_2') else "lora_te_" - module.svdhandler = SVDHandler(maxrank, rank_ratio) - module.svdhandler.network_name = prefix + name.replace(".", "_") - with devices.inference_context(): - module.svdhandler.decompose(module.weight, weights_backup) - t1 = time.time() + # bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + desc, ncols=80, colour='#327fba' + from rich.progress import Progress, TextColumn, BarColumn, TaskProgressColumn, TimeRemainingColumn, TimeElapsedColumn + with Progress(TextColumn('[cyan]LoRA extract'), BarColumn(), TaskProgressColumn(), TimeRemainingColumn(), TimeElapsedColumn(), TextColumn('[cyan]{task.description}'), console=shared.console) as progress: - if 'te' in modules and getattr(shared.sd_model, 'text_encoder_2', None) is not None: - yield "LoRA extract: extracting TE-2" - for name, module in shared.sd_model.text_encoder_2.named_modules(): - weights_backup = getattr(module, "network_weights_backup", None) - if weights_backup is None or getattr(module, "network_current_names", None) is None: - continue - module.svdhandler = SVDHandler(maxrank, rank_ratio) - module.svdhandler.network_name = "lora_te2_" + name.replace(".", "_") - with devices.inference_context(): - module.svdhandler.decompose(module.weight, weights_backup) - t2 = time.time() - - if 'unet' in modules and getattr(shared.sd_model, 'unet', None) is not None: - yield "LoRA extract: extracting UNet" - for name, module in shared.sd_model.unet.named_modules(): - weights_backup = getattr(module, "network_weights_backup", None) - if weights_backup is None or getattr(module, "network_current_names", None) is None: - continue - module.svdhandler = SVDHandler(maxrank, rank_ratio) - module.svdhandler.network_name = "lora_unet_" + name.replace(".", "_") - with devices.inference_context(): - module.svdhandler.decompose(module.weight, weights_backup) - t3 = time.time() - - # TODO: Handle quant for Flux - # if 'te' in modules and getattr(shared.sd_model, 'transformer', None) is not None: - # for name, module in shared.sd_model.transformer.named_modules(): - # if "norm" in name and "linear" not in name: - # continue - # weights_backup = getattr(module, "network_weights_backup", None) - # if weights_backup is None: - # continue - # module.svdhandler = SVDHandler() - # module.svdhandler.network_name = "lora_transformer_" + name.replace(".", "_") - # module.svdhandler.decompose(module.weight, weights_backup) - # module.svdhandler.findrank(rank, rank_ratio) - - lora_state_dict = {} - for sub in ['text_encoder', 'text_encoder_2', 'unet', 'transformer']: - submodel = getattr(shared.sd_model, sub, None) - if submodel is not None: - yield f"LoRA extract: creating {sub}" - for _name, module in submodel.named_modules(): - if not hasattr(module, "svdhandler"): + if 'te' in modules and getattr(shared.sd_model, 'text_encoder', None) is not None: + modules = shared.sd_model.text_encoder.named_modules() + task = progress.add_task(description="te1 decompose", total=len(list(modules))) + for name, module in shared.sd_model.text_encoder.named_modules(): + progress.update(task, advance=1) + weights_backup = getattr(module, "network_weights_backup", None) + if weights_backup is None or getattr(module, "network_current_names", None) is None: continue - lora_state_dict.update(module.svdhandler.makeweights()) - del module.svdhandler - shared.log.debug('LoRA extract: create done') - t4 = time.time() + prefix = "lora_te1_" if hasattr(shared.sd_model, 'text_encoder_2') else "lora_te_" + module.svdhandler = SVDHandler(maxrank, rank_ratio) + module.svdhandler.network_name = prefix + name.replace(".", "_") + with devices.inference_context(): + module.svdhandler.decompose(module.weight, weights_backup) + progress.remove_task(task) + t1 = time.time() + + if 'te' in modules and getattr(shared.sd_model, 'text_encoder_2', None) is not None: + modules = shared.sd_model.text_encoder_2.named_modules() + task = progress.add_task(description="te2 decompose", total=len(list(modules))) + for name, module in shared.sd_model.text_encoder_2.named_modules(): + progress.update(task, advance=1) + weights_backup = getattr(module, "network_weights_backup", None) + if weights_backup is None or getattr(module, "network_current_names", None) is None: + continue + module.svdhandler = SVDHandler(maxrank, rank_ratio) + module.svdhandler.network_name = "lora_te2_" + name.replace(".", "_") + with devices.inference_context(): + module.svdhandler.decompose(module.weight, weights_backup) + progress.remove_task(task) + t2 = time.time() + + if 'unet' in modules and getattr(shared.sd_model, 'unet', None) is not None: + modules = shared.sd_model.unet.named_modules() + task = progress.add_task(description="unet decompose", total=len(list(modules))) + for name, module in shared.sd_model.unet.named_modules(): + progress.update(task, advance=1) + weights_backup = getattr(module, "network_weights_backup", None) + if weights_backup is None or getattr(module, "network_current_names", None) is None: + continue + module.svdhandler = SVDHandler(maxrank, rank_ratio) + module.svdhandler.network_name = "lora_unet_" + name.replace(".", "_") + with devices.inference_context(): + module.svdhandler.decompose(module.weight, weights_backup) + progress.remove_task(task) + t3 = time.time() + + # TODO: Handle quant for Flux + # if 'te' in modules and getattr(shared.sd_model, 'transformer', None) is not None: + # for name, module in shared.sd_model.transformer.named_modules(): + # if "norm" in name and "linear" not in name: + # continue + # weights_backup = getattr(module, "network_weights_backup", None) + # if weights_backup is None: + # continue + # module.svdhandler = SVDHandler() + # module.svdhandler.network_name = "lora_transformer_" + name.replace(".", "_") + # module.svdhandler.decompose(module.weight, weights_backup) + # module.svdhandler.findrank(rank, rank_ratio) + + lora_state_dict = {} + for sub in ['text_encoder', 'text_encoder_2', 'unet', 'transformer']: + submodel = getattr(shared.sd_model, sub, None) + if submodel is not None: + modules = submodel.named_modules() + task = progress.add_task(description=f"{sub} exctract", total=len(list(modules))) + for _name, module in submodel.named_modules(): + progress.update(task, advance=1) + if not hasattr(module, "svdhandler"): + continue + lora_state_dict.update(module.svdhandler.makeweights()) + del module.svdhandler + progress.remove_task(task) + t4 = time.time() if not os.path.isabs(fn): fn = os.path.join(shared.cmd_opts.lora_dir, fn) @@ -200,7 +217,6 @@ def make_lora(fn, maxrank, auto_rank, rank_ratio, modules, overwrite): fn += '.safetensors' if os.path.exists(fn): if overwrite: - shared.log.warning(f'LoRA extract: fn="{fn}" overwriting existing file') os.remove(fn) else: msg = f'LoRA extract: fn="{fn}" file exists' @@ -219,9 +235,9 @@ def make_lora(fn, maxrank, auto_rank, rank_ratio, modules, overwrite): yield msg return t5 = time.time() - shared.log.debug(f'LoRA extract: te1={t1-t0:.2f} te2={t2-t1:.2f} unet={t3-t2:.2f} save={t5-t4:.2f}') + shared.log.debug(f'LoRA extract: time={t5-t0:.2f} te1={t1-t0:.2f} te2={t2-t1:.2f} unet={t3-t2:.2f} save={t5-t4:.2f}') keys = list(lora_state_dict.keys()) - msg = f'LoRA extract: fn="{fn}" keys={len(keys)} time={t5-t0:.2f}' + msg = f'LoRA extract: fn="{fn}" keys={len(keys)}' shared.log.info(msg) yield msg @@ -232,7 +248,7 @@ def create_ui(): with gr.Tab(label="Extract LoRA"): with gr.Row(): - loaded = gr.Textbox(value="Press refresh to query loaded LoRA", label="Loaded LoRA", interactive=False) + loaded = gr.Textbox(placeholder="Press refresh to query loaded LoRA", label="Loaded LoRA", interactive=False) create_refresh_button(loaded, lambda: None, lambda: {'value': loaded_lora_str()}, "testid") with gr.Group(): with gr.Row(): @@ -249,4 +265,8 @@ def create_ui(): status = gr.HTML(value="", show_label=False) auto_rank.change(fn=lambda x: gr_show(x), inputs=[auto_rank], outputs=[rank_ratio]) - extract.click(fn=make_lora, inputs=[filename, rank, auto_rank, rank_ratio, modules, overwrite], outputs=[status]) + extract.click( + fn=wrap_gradio_gpu_call(make_lora, extra_outputs=[]), + inputs=[filename, rank, auto_rank, rank_ratio, modules, overwrite], + outputs=[status] + )