diff --git a/CHANGELOG.md b/CHANGELOG.md index d752952d5..0fb792165 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -171,11 +171,12 @@ For details, see [ChangeLog](https://github.com/vladmandic/automatic/blob/master - remove legacy model configs: `/configs/*.yaml` - remove legacy submodule: `/modules/k-diffusion` - remove legacy hypernetworks support: `/modules/hypernetworks` - - remove legacy lora support: `/extensions-builtin/Lora` - - remove legacy clip/blip interrogate module + - remove legacy lora support: `/extensions-builtin/Lora` + - remove legacy clip/blip interrogate module - remove modern-ui remove `only-original` vs `only-diffusers` code paths - refactor control processing and separate preprocessing and image save ops - refactor modernui layouts to rely on accordions more than individual controls + - refactore pipeline apply/unapply optional components & features - split monolithic `shared.py` - cleanup `/modules`: move pipeline loaders to `/pipelines` root - cleanup `/modules`: move code folders used by pipelines to `/pipelines/` folder diff --git a/modules/control/run.py b/modules/control/run.py index 0db88891e..f66df8399 100644 --- a/modules/control/run.py +++ b/modules/control/run.py @@ -33,8 +33,9 @@ def restore_pipeline(): if instance is not None and hasattr(instance, 'restore'): instance.restore() if (original_pipeline is not None) and (original_pipeline.__class__.__name__ != shared.sd_model.__class__.__name__): - fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access - debug_log(f'Control restored pipeline: class={shared.sd_model.__class__.__name__} to={original_pipeline.__class__.__name__} fn={fn}') + if debug: + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + shared.log.trace(f'Control restored pipeline: class={shared.sd_model.__class__.__name__} to={original_pipeline.__class__.__name__} fn={fn}') shared.sd_model = original_pipeline pipe = None instance = None diff --git a/modules/control/units/controlnet.py b/modules/control/units/controlnet.py index e5a0e1e90..5f3216ba0 100644 --- a/modules/control/units/controlnet.py +++ b/modules/control/units/controlnet.py @@ -378,6 +378,7 @@ class ControlNetPipeline(): unet=pipeline.unet, scheduler=pipeline.scheduler, feature_extractor=getattr(pipeline, 'feature_extractor', None), + image_encoder=getattr(pipeline, 'image_encoder', None), controlnet=controlnets, # can be a list ) elif detect.is_f1(pipeline) and len(controlnets) > 0: @@ -415,6 +416,7 @@ class ControlNetPipeline(): unet=pipeline.unet, scheduler=pipeline.scheduler, feature_extractor=getattr(pipeline, 'feature_extractor', None), + image_encoder=getattr(pipeline, 'image_encoder', None), requires_safety_checker=False, safety_checker=None, controlnet=controlnets, # can be a list diff --git a/modules/devices.py b/modules/devices.py index ed8d91a47..8f5dcbec3 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -391,9 +391,10 @@ def set_cudnn_params(): torch.use_deterministic_algorithms(opts.cudnn_deterministic) if opts.cudnn_deterministic: os.environ.setdefault('CUBLAS_WORKSPACE_CONFIG', ':4096:8') + log.debug('Torch cuDNN: deterministic=True') torch.backends.cudnn.benchmark = opts.cudnn_benchmark if opts.cudnn_benchmark: - log.debug('Torch cuDNN: enable benchmark') + log.debug('Torch cuDNN: benchmark=True') torch.backends.cudnn.benchmark_limit = opts.cudnn_benchmark_limit torch.backends.cudnn.allow_tf32 = True except Exception as e: diff --git a/modules/processing.py b/modules/processing.py index 89f702d0f..3ea931d9b 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -3,12 +3,11 @@ import json import time import numpy as np from PIL import Image, ImageOps -from modules import shared, devices, errors, images, scripts_manager, memstats, script_callbacks, extra_networks, detailer, sd_models, sd_checkpoint, sd_vae, processing_helpers, timer, face_restoration, token_merge +from modules import shared, devices, errors, images, scripts_manager, memstats, script_callbacks, extra_networks, detailer, sd_models, sd_checkpoint, sd_vae, processing_helpers, timer, face_restoration from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet from modules.processing_class import StableDiffusionProcessing, StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img, StableDiffusionProcessingControl, StableDiffusionProcessingVideo # pylint: disable=unused-import from modules.processing_info import create_infotext from modules.modeldata import model_data -from modules import pag, cfgzero opt_C = 4 @@ -158,12 +157,6 @@ def process_images(p: StableDiffusionProcessing) -> Processed: shared.prompt_styles.apply_styles_to_extra(p) shared.prompt_styles.extract_comments(p) - if 'Model' not in shared.opts.cuda_compile: - token_merge.apply_token_merging(p.sd_model) - from modules import sd_hijack_freeu, para_attention, teacache - sd_hijack_freeu.apply_freeu(p) - para_attention.apply_first_block_cache() - teacache.apply_teacache(p) if p.width is not None: p.width = 8 * int(p.width / 8) @@ -205,11 +198,6 @@ def process_images(p: StableDiffusionProcessing) -> Processed: processed = process_images_inner(p) finally: - pag.unapply() - cfgzero.unapply() - if shared.opts.cuda_compile_backend == 'none': - token_merge.remove_token_merging(p.sd_model) - script_callbacks.after_process_callback(p) if p.override_settings_restore_afterwards: # restore opts to original state @@ -284,8 +272,6 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: debug(f'Processing inner: args={vars(p)}') for n in range(p.n_iter): shared.state.batch_no = n + 1 - pag.apply(p) - cfgzero.apply(p) debug(f'Processing inner: iteration={n+1}/{p.n_iter}') p.iteration = n if shared.state.skipped: @@ -296,8 +282,6 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: shared.log.debug(f'Process interrupted: {n+1}/{p.n_iter}') break - from modules import ipadapter - ipadapter.apply(shared.sd_model, p) if not hasattr(p, 'keep_prompts'): p.prompts = p.all_prompts[n * p.batch_size:(n+1) * p.batch_size] p.negative_prompts = p.all_negative_prompts[n * p.batch_size:(n+1) * p.batch_size] @@ -447,9 +431,6 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: if shared.opts.grid_save: images.save_image(grid, p.outpath_grids, "", p.all_seeds[0], p.all_prompts[0], shared.opts.grid_format, info=grid_info, p=p, grid=True, suffix="-grid") # main save grid - from modules import ipadapter - ipadapter.unapply(shared.sd_model, unload=getattr(p, 'ip_adapter_unload', False)) - if shared.opts.include_mask: if shared.opts.mask_apply_overlay and p.overlay_images is not None and len(p.overlay_images) > 0: p.image_mask = create_binary_mask(p.overlay_images[0]) diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 248d32d8b..81118958e 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -5,7 +5,7 @@ import numpy as np import torch import torchvision.transforms.functional as TF from PIL import Image -from modules import shared, devices, processing, sd_models, errors, sd_hijack_hypertile, processing_vae, sd_models_compile, hidiffusion, timer, modelstats, extra_networks, ras, transformer_cache +from modules import shared, devices, processing, sd_models, errors, sd_hijack_hypertile, processing_vae, sd_models_compile, timer, modelstats, extra_networks from modules.processing_helpers import resize_hires, calculate_base_steps, calculate_hires_steps, calculate_refiner_steps, save_intermediate, update_sampler, is_txt2img, is_refiner_enabled, get_job_name from modules.processing_args import set_pipeline_args from modules.onnx_impl import preprocess_pipeline as preprocess_onnx_pipeline, check_parameters_changed as olive_check_parameters_changed @@ -52,6 +52,54 @@ def restore_state(p: processing.StableDiffusionProcessing): return p +def process_pre(p: processing.StableDiffusionProcessing): + from modules import ipadapter, sd_hijack_freeu, para_attention, teacache, hidiffusion, ras, pag, cfgzero, transformer_cache, token_merge + + # apply-unapply + try: + sd_models_compile.check_deepcache(enable=True) + ipadapter.apply(shared.sd_model, p) + token_merge.apply_token_merging(p.sd_model) + hidiffusion.apply(p, shared.sd_model_type) + ras.apply(shared.sd_model, p) + pag.apply(p) + cfgzero.apply(p) + + # apply-only + sd_hijack_freeu.apply_freeu(p) + transformer_cache.set_cache() + para_attention.apply_first_block_cache() + teacache.apply_teacache(p) + except Exception as e: + shared.log.error(f'Processing apply: {e}') + errors.display(e, 'apply') + + if hasattr(shared.sd_model, 'unet'): + sd_models.move_model(shared.sd_model.unet, devices.device) + if hasattr(shared.sd_model, 'transformer'): + sd_models.move_model(shared.sd_model.transformer, devices.device) + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + sd_models.move_model(shared.sd_model, devices.device) + timer.process.record('pre') + + +def process_post(p: processing.StableDiffusionProcessing): + from modules import ipadapter, hidiffusion, ras, pag, cfgzero, token_merge + + try: + sd_models_compile.check_deepcache(enable=False) + ipadapter.unapply(shared.sd_model, unload=getattr(p, 'ip_adapter_unload', False)) + token_merge.remove_token_merging(p.sd_model) + hidiffusion.unapply() + ras.unapply(shared.sd_model) + pag.unapply() + cfgzero.unapply() + except Exception as e: + shared.log.error(f'Processing unapply: {e}') + errors.display(e, 'unapply') + timer.process.record('post') + + def process_base(p: processing.StableDiffusionProcessing): txt2img = is_txt2img() use_refiner_start = is_refiner_enabled(p) and (not p.is_hr_pass) @@ -60,6 +108,7 @@ def process_base(p: processing.StableDiffusionProcessing): shared.sd_model = update_pipeline(shared.sd_model, p) update_sampler(p, shared.sd_model) timer.process.record('prepare') + process_pre(p) base_args = set_pipeline_args( p=p, model=shared.sd_model, @@ -86,18 +135,8 @@ def process_base(p: processing.StableDiffusionProcessing): output = None try: t0 = time.time() - sd_models_compile.check_deepcache(enable=True) - transformer_cache.set_cache() - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - sd_models.move_model(shared.sd_model, devices.device) - if hasattr(shared.sd_model, 'unet'): - sd_models.move_model(shared.sd_model.unet, devices.device) - if hasattr(shared.sd_model, 'transformer'): - sd_models.move_model(shared.sd_model.transformer, devices.device) extra_networks.activate(p, exclude=['text_encoder', 'text_encoder_2', 'text_encoder_3']) - hidiffusion.apply(p, shared.sd_model_type) - ras.apply(shared.sd_model, p) - timer.process.record('move') + if hasattr(shared.sd_model, 'tgate') and getattr(p, 'gate_step', -1) > 0: base_args['gate_step'] = p.gate_step output = shared.sd_model.tgate(**base_args) # pylint: disable=not-callable @@ -143,9 +182,7 @@ def process_base(p: processing.StableDiffusionProcessing): errors.display(e, 'Processing') modelstats.analyze() finally: - ras.unapply(shared.sd_model) - hidiffusion.unapply() - sd_models_compile.check_deepcache(enable=False) + process_post(p) shared.state.nextjob() return output @@ -206,6 +243,7 @@ def process_hires(p: processing.StableDiffusionProcessing, output): orig_denoise = p.denoising_strength p.denoising_strength = strength orig_image = p.task_args.pop('image', None) # remove image override from hires + process_pre(p) hires_args = set_pipeline_args( p=p, model=shared.sd_model, @@ -226,15 +264,8 @@ def process_hires(p: processing.StableDiffusionProcessing, output): hires_steps = hires_args.get('prior_num_inference_steps', None) or p.hr_second_pass_steps or hires_args.get('num_inference_steps', None) shared.state.update(get_job_name(p, shared.sd_model), hires_steps, 1) try: - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - sd_models.move_model(shared.sd_model, devices.device) - if hasattr(shared.sd_model, 'unet'): - sd_models.move_model(shared.sd_model.unet, devices.device) - if hasattr(shared.sd_model, 'transformer'): - sd_models.move_model(shared.sd_model.transformer, devices.device) if 'base' in p.skip: extra_networks.activate(p) - sd_models_compile.check_deepcache(enable=True) output = shared.sd_model(**hires_args) # pylint: disable=not-callable if isinstance(output, dict): output = SimpleNamespace(**output) @@ -249,6 +280,8 @@ def process_hires(p: processing.StableDiffusionProcessing, output): shared.log.error(f'Processing step=hires: args={hires_args} {e}') errors.display(e, 'Processing') modelstats.analyze() + finally: + process_post(p) if orig_image is not None: p.task_args['image'] = orig_image p.denoising_strength = orig_denoise diff --git a/modules/sd_models.py b/modules/sd_models.py index d8ef24a3e..774b903a5 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -812,6 +812,55 @@ def clean_diffuser_pipe(pipe): pipe.register_to_config(**internal_dict) +def backup_pipe_components(pipe): + if pipe is None: + return {} + return { + 'sd_checkpoint_info': getattr(pipe, "sd_checkpoint_info", None), + 'sd_model_checkpoint': getattr(pipe, "sd_model_checkpoint", None), + 'embedding_db': getattr(pipe, "embedding_db", None), + 'loaded_loras': getattr(pipe, "loaded_loras", None), + 'sd_model_hash': getattr(pipe, "sd_model_hash", None), + 'has_accelerate': getattr(pipe, "has_accelerate", None), + 'current_attn_name': getattr(pipe, "current_attn_name", None), + 'default_scheduler': getattr(pipe, "default_scheduler", None), + 'image_encoder': getattr(pipe, "image_encoder", None), + 'feature_extractor': getattr(pipe, "feature_extractor", None), + 'mask_processor': getattr(pipe, "mask_processor", None), + 'restore_pipeline': getattr(pipe, "restore_pipeline", None), + } + + +def restore_pipe_components(pipe, components): + if pipe is None or components is None: + return + if hasattr(pipe, 'sd_checkpoint_info'): + pipe.sd_checkpoint_info = components['sd_checkpoint_info'] + if hasattr(pipe, 'sd_model_checkpoint'): + pipe.sd_model_checkpoint = components['sd_model_checkpoint'] + if hasattr(pipe, 'embedding_db'): + pipe.embedding_db = components['embedding_db'] + if hasattr(pipe, 'loaded_loras'): + pipe.loaded_loras = components['loaded_loras'] if components['loaded_loras'] is not None else {} + if hasattr(pipe, 'sd_model_hash'): + pipe.sd_model_hash = components['sd_model_hash'] + if hasattr(pipe, 'has_accelerate'): + pipe.has_accelerate = components['has_accelerate'] + if hasattr(pipe, 'current_attn_name'): + pipe.current_attn_name = components['current_attn_name'] + if hasattr(pipe, 'default_scheduler'): + pipe.default_scheduler = components['default_scheduler'] + if hasattr(pipe, 'image_encoder') and components['image_encoder'] is not None: + pipe.image_encoder = components['image_encoder'] + if hasattr(pipe, 'feature_extractor') and components['feature_extractor'] is not None: + pipe.feature_extractor = components['feature_extractor'] + if hasattr(pipe, 'mask_processor') and components['mask_processor'] is not None: + pipe.mask_processor = components['mask_processor'] + if pipe.__class__.__name__ in ['FluxPipeline', 'StableDiffusion3Pipeline']: + pipe.register_modules(image_encoder = components['image_encoder']) + pipe.register_modules(feature_extractor = components['feature_extractor']) + + def set_diffuser_pipe(pipe, new_pipe_type): has_errors = False if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: @@ -840,18 +889,7 @@ def set_diffuser_pipe(pipe, new_pipe_type): if cls == 'StableDiffusionXLPAGPipeline': pipe = switch_pipe(diffusers.StableDiffusionXLPipeline, pipe) - sd_checkpoint_info = getattr(pipe, "sd_checkpoint_info", None) - sd_model_checkpoint = getattr(pipe, "sd_model_checkpoint", None) - embedding_db = getattr(pipe, "embedding_db", None) - loaded_loras = getattr(pipe, "loaded_loras", None) - sd_model_hash = getattr(pipe, "sd_model_hash", None) - has_accelerate = getattr(pipe, "has_accelerate", None) - current_attn_name = getattr(pipe, "current_attn_name", None) - default_scheduler = getattr(pipe, "default_scheduler", None) - image_encoder = getattr(pipe, "image_encoder", None) - feature_extractor = getattr(pipe, "feature_extractor", None) - mask_processor = getattr(pipe, "mask_processor", None) - restore_pipeline = getattr(pipe, "restore_pipeline", None) + components_backup = backup_pipe_components(pipe) if new_pipe is None: if hasattr(pipe, 'config'): # real pipeline which can be auto-switched @@ -889,25 +927,10 @@ def set_diffuser_pipe(pipe, new_pipe_type): if new_pipe is None: return pipe - new_pipe.sd_checkpoint_info = sd_checkpoint_info - new_pipe.sd_model_checkpoint = sd_model_checkpoint - new_pipe.embedding_db = embedding_db - new_pipe.sd_model_hash = sd_model_hash - new_pipe.has_accelerate = has_accelerate - new_pipe.current_attn_name = current_attn_name - new_pipe.default_scheduler = default_scheduler - new_pipe.loaded_loras = loaded_loras if loaded_loras is not None else {} - if image_encoder is not None: - new_pipe.image_encoder = image_encoder - if feature_extractor is not None: - new_pipe.feature_extractor = feature_extractor - if mask_processor is not None: - new_pipe.mask_processor = mask_processor - if restore_pipeline is not None: - new_pipe.restore_pipeline = restore_pipeline - if new_pipe.__class__.__name__ in ['FluxPipeline', 'StableDiffusion3Pipeline']: - new_pipe.register_modules(image_encoder = image_encoder) - new_pipe.register_modules(feature_extractor = feature_extractor) + + restore_pipe_components(new_pipe, components_backup) + components_backup = None # free memory + new_pipe.is_sdxl = getattr(pipe, 'is_sdxl', False) # a1111 compatibility item new_pipe.is_sd2 = getattr(pipe, 'is_sd2', False) new_pipe.is_sd1 = getattr(pipe, 'is_sd1', True) diff --git a/modules/sd_offload.py b/modules/sd_offload.py index 668d4eaaf..46aa02ece 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -248,7 +248,7 @@ class OffloadHook(accelerate.hooks.ModelHook): elif 'bitsandbytes' in str(e): pass else: - shared.log.error(f'Offload: type=balanced op=apply module={module.__name__} {e}') + shared.log.error(f'Offload: type=balanced op=apply module={module.__name__} cls={module.__class__ if inspect.isclass(module) else None} {e}') if os.environ.get('SD_MOVE_DEBUG', None): errors.display(e, f'Offload: type=balanced op=apply module={module.__name__}') return output