update swinir

This commit is contained in:
Vladimir Mandic
2023-06-08 08:16:43 -04:00
parent 6ef3e964a3
commit 727b41a4bd
@@ -3,7 +3,7 @@ import numpy as np
import torch
from PIL import Image
from basicsr.utils.download_util import load_file_from_url
from tqdm import tqdm
from tqdm.rich import tqdm
from swinir_model_arch import SwinIR as net
from swinir_model_arch_v2 import Swin2SR as net2
from modules import modelloader, devices, script_callbacks, shared
@@ -54,8 +54,7 @@ class UpscalerSwinIR(Upscaler):
filename = path
if filename is None or not os.path.exists(filename):
return None
if filename.endswith(".v2.pth"):
model = net2(
model_v2 = net2(
upscale=scale,
in_chans=3,
img_size=64,
@@ -67,29 +66,33 @@ class UpscalerSwinIR(Upscaler):
mlp_ratio=2,
upsampler="nearest+conv",
resi_connection="1conv",
)
params = None
else:
model = net(
upscale=scale,
in_chans=3,
img_size=64,
window_size=8,
img_range=1.0,
depths=[6, 6, 6, 6, 6, 6, 6, 6, 6],
embed_dim=240,
num_heads=[8, 8, 8, 8, 8, 8, 8, 8, 8],
mlp_ratio=2,
upsampler="nearest+conv",
resi_connection="3conv",
)
params = "params_ema"
)
model_v1 = net(
upscale=scale,
in_chans=3,
img_size=64,
window_size=8,
img_range=1.0,
depths=[6, 6, 6, 6, 6, 6, 6, 6, 6],
embed_dim=240,
num_heads=[8, 8, 8, 8, 8, 8, 8, 8, 8],
mlp_ratio=2,
upsampler="nearest+conv",
resi_connection="3conv",
)
pretrained_model = torch.load(filename)
if params is not None:
model.load_state_dict(pretrained_model[params], strict=True)
else:
model.load_state_dict(pretrained_model, strict=True)
for model in [model_v1, model_v2]:
for param in ["params_ema", "params", None]:
try:
if param is not None:
model.load_state_dict(pretrained_model[param], strict=True)
else:
model.load_state_dict(pretrained_model, strict=True)
shared.log.info(f'Loaded SwinIR model: {filename} param={param}')
return model
except Exception:
pass
shared.log.error(f'Could not determine SwinIR model parameters: {filename}')
return model
@@ -140,7 +143,7 @@ def inference(img, model, tile, tile_overlap, window_size, scale):
E = torch.zeros(b, c, h * sf, w * sf, dtype=devices.dtype, device=device_swinir).type_as(img)
W = torch.zeros_like(E, dtype=devices.dtype, device=device_swinir)
with tqdm(total=len(h_idx_list) * len(w_idx_list), desc="SwinIR tiles") as pbar:
with tqdm(total=len(h_idx_list) * len(w_idx_list), desc="Upscaling SwinIR") as pbar:
for h_idx in h_idx_list:
if state.interrupted or state.skipped:
break