mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
upscaler caching and ti model detection
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user