upscaler caching and ti model detection

This commit is contained in:
Vladimir Mandic
2023-09-27 08:49:04 -04:00
parent 73b6d4f57c
commit ef0e8a5161
12 changed files with 140 additions and 106 deletions
-20
View File
@@ -45,7 +45,6 @@ def setup_model(dirname):
self.cmd_dir = dirname
def create_models(self):
if self.net is not None and self.face_helper is not None:
self.net.to(devices.device_codeformer)
return self.net, self.face_helper
@@ -60,14 +59,11 @@ def setup_model(dirname):
net.load_state_dict(checkpoint)
net.eval()
shared.log.info(f"Model loaded: type=CodeFormer model={ckpt_path}")
if hasattr(retinaface, 'device'):
retinaface.device = devices.device_codeformer
face_helper = FaceRestoreHelper(1, face_size=512, crop_ratio=(1, 1), det_model='retinaface_resnet50', save_ext='png', use_parse=True, device=devices.device_codeformer)
self.net = net
self.face_helper = face_helper
return net, face_helper
def send_model_to(self, device):
@@ -77,25 +73,19 @@ def setup_model(dirname):
def restore(self, np_image, w=None):
np_image = np_image[:, :, ::-1]
original_resolution = np_image.shape[0:2]
self.create_models()
if self.net is None or self.face_helper is None:
return np_image
self.send_model_to(devices.device_codeformer)
self.face_helper.clean_all()
self.face_helper.read_image(np_image)
self.face_helper.get_face_landmarks_5(only_center_face=False, resize=640, eye_dist_threshold=5)
self.face_helper.align_warp_face()
for cropped_face in self.face_helper.cropped_faces:
cropped_face_t = img2tensor(cropped_face / 255., bgr2rgb=True, float32=True)
normalize(cropped_face_t, (0.5, 0.5, 0.5), (0.5, 0.5, 0.5), inplace=True)
cropped_face_t = cropped_face_t.unsqueeze(0).to(devices.device_codeformer)
try:
with devices.inference_context():
output = self.net(cropped_face_t, w=w if w is not None else shared.opts.code_former_weight, adain=True)[0]
@@ -105,33 +95,23 @@ def setup_model(dirname):
except Exception as e:
shared.log.error(f'CodeForomer error: {e}')
restored_face = tensor2img(cropped_face_t, rgb2bgr=True, min_max=(-1, 1))
restored_face = restored_face.astype('uint8')
self.face_helper.add_restored_face(restored_face)
self.face_helper.get_inverse_affine(None)
restored_img = self.face_helper.paste_faces_to_input_image()
restored_img = restored_img[:, :, ::-1]
if original_resolution != restored_img.shape[0:2]:
restored_img = cv2.resize(restored_img, (0, 0), fx=original_resolution[1]/restored_img.shape[1], fy=original_resolution[0]/restored_img.shape[0], interpolation=cv2.INTER_LINEAR)
self.face_helper.clean_all()
if shared.opts.face_restoration_unload:
self.send_model_to(devices.cpu)
return restored_img
global have_codeformer # pylint: disable=global-statement
have_codeformer = True
global codeformer # pylint: disable=global-statement
codeformer = FaceRestorerCodeFormer(dirname)
shared.face_restorers.append(codeformer)
except Exception as e:
errors.display(e, 'codeformer')
# sys.path = stored_sys_path
+15 -8
View File
@@ -4,7 +4,7 @@ from PIL import Image
from rich.progress import Progress, TextColumn, BarColumn, TaskProgressColumn, TimeRemainingColumn, TimeElapsedColumn
import modules.postprocess.esrgan_model_arch as arch
from modules import images, devices
from modules.upscaler import Upscaler
from modules.upscaler import Upscaler, UpscalerData
from modules.shared import opts, log, console
@@ -122,7 +122,7 @@ class UpscalerESRGAN(Upscaler):
self.user_path = dirname
super().__init__()
self.scalers = self.find_scalers()
self.models = {}
def do_upscale(self, img, selected_model):
model = self.load_model(selected_model)
@@ -130,12 +130,19 @@ class UpscalerESRGAN(Upscaler):
return img
model.to(devices.device_esrgan)
img = esrgan_upscale(model, img)
if opts.upscaler_unload and selected_model in self.models:
del self.models[selected_model]
log.debug(f"Upscaler unloaded: type={self.name} model={selected_model}")
devices.torch_gc(force=True)
return img
def load_model(self, path: str):
info = self.find_model(path)
info: UpscalerData = self.find_model(path)
if info is None:
return
if self.models.get(info.local_data_path, None) is not None:
log.debug(f"Upscaler cached: type={self.name} model={info.local_data_path}")
return self.models[info.local_data_path]
state_dict = torch.load(info.local_data_path, map_location='cpu' if devices.device_esrgan.type == 'mps' else None)
log.info(f"Upscaler loaded: type={self.name} model={info.local_data_path}")
@@ -147,7 +154,9 @@ class UpscalerESRGAN(Upscaler):
model = arch.SRVGGNetCompact(num_in_ch=3, num_out_ch=3, num_feat=64, num_conv=num_conv, upscale=4, act_type='prelu')
model.load_state_dict(state_dict)
model.eval()
return model
self.models[info.local_data_path] = model
return self.models[info.local_data_path]
if "body.0.rdb1.conv1.weight" in state_dict and "conv_first.weight" in state_dict:
nb = 6 if "RealESRGAN_x4plus_anime_6B" in info.local_data_path else 23
state_dict = resrgan2normal(state_dict, nb)
@@ -155,14 +164,12 @@ class UpscalerESRGAN(Upscaler):
state_dict = mod2normal(state_dict)
elif "model.0.weight" not in state_dict:
raise TypeError("The file is not a recognized ESRGAN model.")
in_nc, out_nc, nf, nb, plus, mscale = infer_params(state_dict)
model = arch.RRDBNet(in_nc=in_nc, out_nc=out_nc, nf=nf, nb=nb, upscale=mscale, plus=plus)
model.load_state_dict(state_dict)
model.eval()
return model
self.models[info.local_data_path] = model
return self.models[info.local_data_path]
def upscale_without_tiling(model, img):
+21 -13
View File
@@ -5,7 +5,7 @@ from basicsr.archs.rrdbnet_arch import RRDBNet
from modules.postprocess.realesrgan_model_arch import SRVGGNetCompact
from modules.upscaler import Upscaler
from modules.shared import opts, device, log
from modules import devices
class UpscalerRealESRGAN(Upscaler):
def __init__(self, dirname):
@@ -13,6 +13,7 @@ class UpscalerRealESRGAN(Upscaler):
self.user_path = dirname
super().__init__()
self.scalers = self.find_scalers()
self.models = {}
for scaler in self.scalers:
if scaler.name == 'RealESRGAN 2x+':
scaler.model = lambda: RRDBNet(num_in_ch=3, num_out_ch=3, num_feat=64, num_block=23, num_grow_ch=32, scale=2)
@@ -41,22 +42,29 @@ class UpscalerRealESRGAN(Upscaler):
except Exception:
log.error("Error importing Real-ESRGAN:")
return img
info = self.find_model(selected_model)
if info is None or not os.path.exists(info.local_data_path):
return img
upsampler = RealESRGANer(
scale=info.scale,
model_path=info.local_data_path,
model=info.model(),
half=not opts.no_half and not opts.upcast_sampling,
tile=opts.ESRGAN_tile,
tile_pad=opts.ESRGAN_tile_overlap,
device=device,
)
if self.models.get(info.local_data_path, None) is not None:
log.debug(f"Upscaler cached: type={self.name} model={info.local_data_path}")
upsampler=self.models[info.local_data_path]
else:
upsampler = RealESRGANer(
name=info.name,
scale=info.scale,
model_path=info.local_data_path,
model=info.model(),
half=not opts.no_half and not opts.upcast_sampling,
tile=opts.ESRGAN_tile,
tile_pad=opts.ESRGAN_tile_overlap,
device=device,
)
self.models[info.local_data_path] = upsampler
upsampled = upsampler.enhance(np.array(img), outscale=info.scale)[0]
if opts.upscaler_unload and info.local_data_path in self.models:
del self.models[info.local_data_path]
log.debug(f"Upscaler unloaded: type={self.name} model={selected_model}")
devices.torch_gc(force=True)
image = Image.fromarray(upsampled)
return image
@@ -29,6 +29,7 @@ class RealESRGANer():
"""
def __init__(self,
name,
scale,
model_path,
dni_weight=None,
@@ -39,6 +40,7 @@ class RealESRGANer():
half=False,
device=None,
gpu_id=None):
self.name = name
self.scale = scale
self.tile_size = tile
self.tile_pad = tile_pad
@@ -63,6 +65,7 @@ class RealESRGANer():
from modules.modelloader import load_file_from_url
model_path = load_file_from_url(url=model_path, model_dir=os.path.join(ROOT_DIR, 'weights'), progress=True, file_name=None)
loadnet = torch.load(model_path, map_location=torch.device('cpu'))
log.info(f"Upscaler loaded: type={self.name} model={model_path}")
# prefer to use params_ema
if 'params_ema' in loadnet:
+19 -8
View File
@@ -14,18 +14,24 @@ class UpscalerSCUNet(Upscaler):
self.user_path = dirname
super().__init__()
self.scalers = self.find_scalers()
self.models = {}
def load_model(self, path: str):
info = self.find_model(path)
if info is None:
return
model = net(in_nc=3, config=[4, 4, 4, 4, 4, 4, 4], dim=64)
model.load_state_dict(torch.load(info.local_data_path), strict=True)
model.eval()
log.info(f"Upscaler loaded: type={self.name} model={info.local_data_path}")
for _, v in model.named_parameters():
v.requires_grad = False
model = model.to(device)
if self.models.get(info.local_data_path, None) is not None:
log.debug(f"Upscaler cached: type={self.name} model={info.local_data_path}")
model=self.models[info.local_data_path]
else:
model = net(in_nc=3, config=[4, 4, 4, 4, 4, 4, 4], dim=64)
model.load_state_dict(torch.load(info.local_data_path), strict=True)
model.eval()
log.info(f"Upscaler loaded: type={self.name} model={info.local_data_path}")
for _, v in model.named_parameters():
v.requires_grad = False
model = model.to(device)
self.models[info.local_data_path] = model
return model
@staticmethod
@@ -83,7 +89,12 @@ class UpscalerSCUNet(Upscaler):
devices.torch_gc()
output = np_output.transpose((1, 2, 0)) # CHW to HWC
output = output[:, :, ::-1] # BGR to RGB
return PIL.Image.fromarray((output * 255).astype(np.uint8))
img = PIL.Image.fromarray((output * 255).astype(np.uint8))
if opts.upscaler_unload and selected_file in self.models:
del self.models[selected_file]
log.debug(f"Upscaler unloaded: type={self.name} model={selected_file}")
devices.torch_gc(force=True)
return img
def on_ui_settings():
+15 -6
View File
@@ -19,17 +19,22 @@ class UpscalerSD(Upscaler):
None,
None,
]
self.models = {}
def load_model(self, path: str):
from modules.sd_models import set_diffuser_options
scaler = [x for x in self.scalers if x.data_path == path][0]
if scaler.model is None:
scaler: UpscalerData = [x for x in self.scalers if x.data_path == path][0]
if self.models.get(path, None) is not None:
shared.log.debug(f"Upscaler cached: type={scaler.name} model={path}")
return self.models[path]
else:
devices.set_cuda_params()
scaler.model = diffusers.DiffusionPipeline.from_pretrained(path, cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype)
if hasattr(scaler.model, "set_progress_bar_config"):
scaler.model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + 'Upscale', ncols=80, colour='#327fba')
model = diffusers.DiffusionPipeline.from_pretrained(path, cache_dir=shared.opts.diffusers_dir, torch_dtype=devices.dtype)
if hasattr(model, "set_progress_bar_config"):
model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + 'Upscale', ncols=80, colour='#327fba')
set_diffuser_options(scaler.model, vae=None, op='upscaler')
return scaler.model
self.models[path] = model
return self.models[path]
def callback(self, _step: int, _timestep: int, _latents: torch.FloatTensor):
pass
@@ -61,4 +66,8 @@ class UpscalerSD(Upscaler):
model = model.to(devices.device)
output = model(**args)
image = output.images[0]
if shared.opts.upscaler_unload and selected_model in self.models:
del self.models[selected_model]
shared.log.debug(f"Upscaler unloaded: type={self.name} model={selected_model}")
devices.torch_gc(force=True)
return image
+9 -1
View File
@@ -14,11 +14,15 @@ class UpscalerSwinIR(Upscaler):
self.user_path = dirname
super().__init__()
self.scalers = self.find_scalers()
self.models = {}
def load_model(self, path, scale=4):
info = self.find_model(path)
if info is None:
return
if self.models.get(info.local_data_path, None) is not None:
shared.log.debug(f"Upscaler cached: type={self.name} model={info.local_data_path}")
return self.models[info.local_data_path]
pretrained_model = torch.load(info.local_data_path)
model_v2 = net2(
upscale=scale,
@@ -54,6 +58,7 @@ class UpscalerSwinIR(Upscaler):
else:
model.load_state_dict(pretrained_model, strict=True)
shared.log.info(f"Upscaler loaded: type={self.name} model={info.local_data_path} param={param}")
self.models[info.local_data_path] = model
return model
except Exception as e:
shared.log.error(f'Upscaler invalid parameters: type={self.name} model={info.local_data_path} {e}')
@@ -65,7 +70,10 @@ class UpscalerSwinIR(Upscaler):
return img
model = model.to(shared.device, dtype=devices.dtype)
img = upscale(img, model)
devices.torch_gc()
if shared.opts.upscaler_unload and selected_model in self.models:
del self.models[selected_model]
shared.log.debug(f"Upscaler unloaded: type={self.name} model={selected_model}")
devices.torch_gc(force=True)
return img