From 2190788769c3f299ec1a61b765496dd1b42db87a Mon Sep 17 00:00:00 2001 From: Disty0 Date: Fri, 5 Jan 2024 20:41:33 +0300 Subject: [PATCH] OpenVINO convert unused weight to FakeTensors --- CHANGELOG.md | 1 + modules/intel/openvino/__init__.py | 6 +++--- modules/processing_diffusers.py | 4 ++-- modules/processing_vae.py | 4 ++-- modules/sd_models.py | 7 +++++++ 5 files changed, 15 insertions(+), 7 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index beebef73e..102bc2877 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -77,6 +77,7 @@ And it also includes fixes for all reported issues so far - **4-bit support with NNCF** enable *Compress Model weights with NNCF* from *Compute Settings* and set a 4-bit NNCF mode 4-bit and 8-bit with OpenVINO is CPU only for now + - reduce system memory usage after compile - **Fixes** - ipadapter: allow changing of model/image on-the-fly - ipadapter: fix fallback of cross-attention on unload diff --git a/modules/intel/openvino/__init__.py b/modules/intel/openvino/__init__.py index 74890491b..441a12820 100644 --- a/modules/intel/openvino/__init__.py +++ b/modules/intel/openvino/__init__.py @@ -18,7 +18,7 @@ from types import MappingProxyType from hashlib import sha256 import functools -from modules import shared, devices +from modules import shared, devices, sd_models NNCFNodeName = str def get_node_by_name(self, name: NNCFNodeName) -> nncf.common.graph.NNCFNode: @@ -376,8 +376,8 @@ def openvino_fx(subgraph, example_inputs): example_inputs_reordered.append(example_inputs[idx1]) example_inputs = example_inputs_reordered - # Deleting unused subgraphs doesn't do anything, so we cast it down to fp8 - subgraph = subgraph.to(dtype=torch.float8_e4m3fn) + # Delete unused subgraphs + subgraph = subgraph.apply(sd_models.convert_to_faketensors) devices.torch_gc(force=True) # Model is fully supported and already cached. Run the cached OV model directly. diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 71808bf98..dc6572074 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -350,12 +350,12 @@ def process_diffusers(p: StableDiffusionProcessing): if shared.compiled_model_state.first_pass and op == "base": shared.compiled_model_state.first_pass = False if hasattr(shared.sd_model, "unet"): - shared.sd_model.unet.to(dtype=torch.float8_e4m3fn) + shared.sd_model.unet.apply(sd_models.convert_to_faketensors) devices.torch_gc(force=True) if shared.compiled_model_state.first_pass_refiner and op == "refiner": shared.compiled_model_state.first_pass_refiner = False if hasattr(shared.sd_refiner, "unet"): - shared.sd_refiner.unet.to(dtype=torch.float8_e4m3fn) + shared.sd_refiner.unet.apply(sd_models.convert_to_faketensors) devices.torch_gc(force=True) def update_sampler(sd_model, second_pass=False): diff --git a/modules/processing_vae.py b/modules/processing_vae.py index 58b6b0ab1..7b5fba688 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -2,7 +2,7 @@ import os import time import torch import torchvision.transforms.functional as TF -from modules import shared, devices, sd_vae +from modules import shared, devices, sd_vae, sd_models import modules.taesd.sd_vae_taesd as sd_vae_taesd @@ -58,7 +58,7 @@ def full_vae_decode(latents, model): if shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx" and shared.compiled_model_state.first_pass_vae: shared.compiled_model_state.first_pass_vae = False if hasattr(shared.sd_model, "vae"): - model.vae.to(dtype=torch.float8_e4m3fn) + model.vae.apply(sd_models.convert_to_faketensors) devices.torch_gc(force=True) if shared.opts.diffusers_move_unet and not getattr(model, 'has_accelerate', False) and hasattr(model, 'unet'): diff --git a/modules/sd_models.py b/modules/sd_models.py index 99ee45b1a..f27b82181 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1262,6 +1262,13 @@ def reload_model_weights(sd_model=None, info=None, reuse_dict=False, op='model') return sd_model +def convert_to_faketensors(tensor): + fake = torch._subclasses.fake_tensor.FakeTensorMode() + if hasattr(tensor, "weight"): + tensor.weight = torch.nn.Parameter(fake.from_tensor(tensor.weight)) + return tensor + + def disable_offload(sd_model): from accelerate.hooks import remove_hook_from_module if not getattr(sd_model, 'has_accelerate', False):