update_inference
This commit is contained in:
+25
-36
@@ -29,7 +29,6 @@ def process(
|
||||
t_min: float,
|
||||
strength: float,
|
||||
color_fix_type: str,
|
||||
disable_preprocess_model: bool,
|
||||
cond_fn: Optional[MSEGuidance],
|
||||
tiled: bool,
|
||||
tile_size: int,
|
||||
@@ -46,7 +45,6 @@ def process(
|
||||
t_min (float):
|
||||
strength (float): Control strength. Set to 1.0 during training.
|
||||
color_fix_type (str): Type of color correction for samples.
|
||||
disable_preprocess_model (bool): If specified, preprocess model (SwinIR) will not be used.
|
||||
cond_fn (Guidance | None): Guidance function that returns gradient to guide the predicted x_0.
|
||||
tiled (bool): If specified, a patch-based sampling strategy will be used for sampling.
|
||||
tile_size (int): Size of patch.
|
||||
@@ -54,17 +52,12 @@ def process(
|
||||
|
||||
Returns:
|
||||
preds (List[np.ndarray]): Restoration results (HWC, RGB, range in [0, 255]).
|
||||
stage1_preds (List[np.ndarray]): Outputs of preprocess model (HWC, RGB, range in [0, 255]).
|
||||
If `disable_preprocess_model` is specified, then preprocess model's outputs is the same
|
||||
as low-quality inputs.
|
||||
"""
|
||||
n_samples = len(control_imgs)
|
||||
sampler = SpacedSampler(model, var_type="fixed_small")
|
||||
control = torch.tensor(np.stack(control_imgs) / 255.0, dtype=torch.float32, device=model.device).clamp_(0, 1)
|
||||
control = einops.rearrange(control, "n h w c -> n c h w").contiguous()
|
||||
|
||||
if not disable_preprocess_model:
|
||||
control = model.preprocess_model(control)
|
||||
model.control_scales = [strength] * 13
|
||||
|
||||
if cond_fn is not None:
|
||||
@@ -73,8 +66,6 @@ def process(
|
||||
height, width = control.size(-2), control.size(-1)
|
||||
shape = (n_samples, 4, height // 8, width // 8)
|
||||
x_T = torch.randn(shape, device=model.device, dtype=torch.float32)
|
||||
init_latent = model.encode_first_stage(control)
|
||||
init_latent = model.get_first_stage_encoding(init_latent)
|
||||
if not tiled:
|
||||
# samples = sampler.sample_ccsr_stage1(
|
||||
# steps=steps, t_max=t_max, shape=shape, cond_img=control,
|
||||
@@ -96,36 +87,34 @@ def process(
|
||||
cfg_scale=1.0, cond_fn=cond_fn,
|
||||
color_fix_type=color_fix_type
|
||||
)
|
||||
|
||||
x_samples = samples.clamp(0, 1)
|
||||
x_samples = (einops.rearrange(x_samples, "b c h w -> b h w c") * 255).cpu().numpy().clip(0, 255).astype(np.uint8)
|
||||
control = (einops.rearrange(control, "b c h w -> b h w c") * 255).cpu().numpy().clip(0, 255).astype(np.uint8)
|
||||
|
||||
preds = [x_samples[i] for i in range(n_samples)]
|
||||
stage1_preds = [control[i] for i in range(n_samples)]
|
||||
|
||||
return preds, stage1_preds
|
||||
return preds
|
||||
|
||||
|
||||
def parse_args() -> Namespace:
|
||||
parser = ArgumentParser()
|
||||
|
||||
# TODO: add help info for these options
|
||||
parser.add_argument("--ckpt", type=str, help="full checkpoint path", default='/home/notebook/data/group/SunLingchen/code/CCSR/CCSR_weights/step=59.ckpt')
|
||||
parser.add_argument("--ckpt", type=str, help="full checkpoint path",
|
||||
default='weights/real-world_ccsr.ckpt')
|
||||
parser.add_argument("--config", type=str, help="model config path", default='configs/model/ccsr_stage2.yaml')
|
||||
|
||||
parser.add_argument("--input", type=str, default='inputs/real47')
|
||||
parser.add_argument("--input", type=str, default='preset/test_datasets')
|
||||
parser.add_argument("--steps", type=int, default=45)
|
||||
parser.add_argument("--sr_scale", type=float, default=4)
|
||||
parser.add_argument("--repeat_times", type=int, default=1)
|
||||
parser.add_argument("--disable_preprocess_model", action="store_true")
|
||||
|
||||
# patch-based sampling
|
||||
# patch-based sampling (tiling settings)
|
||||
parser.add_argument("--tiled", action="store_true")
|
||||
parser.add_argument("--tile_size", type=int, default=512)
|
||||
parser.add_argument("--tile_stride", type=int, default=256)
|
||||
parser.add_argument("--tile_size", type=int, default=512) # image size
|
||||
parser.add_argument("--tile_stride", type=int, default=256) # image size
|
||||
|
||||
parser.add_argument("--color_fix_type", type=str, default="adain", choices=["wavelet", "adain", "none"])
|
||||
parser.add_argument("--output", type=str,default="experiments/output")
|
||||
parser.add_argument("--output", type=str, default="experiments/test")
|
||||
parser.add_argument("--t_max", type=float, default=0.6667)
|
||||
parser.add_argument("--t_min", type=float, default=0.3333)
|
||||
parser.add_argument("--show_lq", action="store_true")
|
||||
@@ -136,6 +125,7 @@ def parse_args() -> Namespace:
|
||||
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def check_device(device):
|
||||
if device == "cuda":
|
||||
# check if CUDA is available
|
||||
@@ -160,6 +150,7 @@ def check_device(device):
|
||||
print(f'using device {device}')
|
||||
return device
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
pl.seed_everything(args.seed)
|
||||
@@ -169,13 +160,7 @@ def main() -> None:
|
||||
model: ControlLDM = instantiate_from_config(OmegaConf.load(args.config))
|
||||
load_state_dict(model, torch.load(args.ckpt, map_location="cpu"), strict=True)
|
||||
# reload preprocess model if specified
|
||||
args.reload_swinir = False
|
||||
args.disable_preprocess_model = True
|
||||
if args.reload_swinir:
|
||||
if not hasattr(model, "preprocess_model"):
|
||||
raise ValueError(f"model don't have a preprocess model.")
|
||||
print(f"reload swinir model from {args.swinir_ckpt}")
|
||||
load_state_dict(model.preprocess_model, torch.load(args.swinir_ckpt, map_location="cpu"), strict=True)
|
||||
|
||||
model.freeze()
|
||||
model.to(args.device)
|
||||
|
||||
@@ -192,12 +177,17 @@ def main() -> None:
|
||||
lq_resized = auto_resize(lq, 512)
|
||||
else:
|
||||
lq_resized = auto_resize(lq, args.tile_size)
|
||||
x = pad(np.array(lq_resized), scale=64)
|
||||
|
||||
x = lq_resized.resize(
|
||||
tuple(s // 64 * 64 for s in lq_resized.size), Image.LANCZOS
|
||||
)
|
||||
x = np.array(x)
|
||||
# x = pad(np.array(lq_resized), scale=64)
|
||||
|
||||
for i in range(args.repeat_times):
|
||||
save_path = os.path.join(args.output, os.path.relpath(file_path, args.input))
|
||||
parent_path, stem, _ = get_file_name_parts(save_path)
|
||||
save_path_now = os.path.join(parent_path, 'sample'+str(i))
|
||||
save_path_now = os.path.join(parent_path, 'sample' + str(i))
|
||||
|
||||
save_path = os.path.join(save_path_now, f"{stem}.png")
|
||||
if os.path.exists(save_path):
|
||||
@@ -206,36 +196,35 @@ def main() -> None:
|
||||
continue
|
||||
else:
|
||||
raise RuntimeError(f"{save_path} already exist")
|
||||
# os.makedirs(parent_path, exist_ok=True)
|
||||
|
||||
os.makedirs(save_path_now, exist_ok=True)
|
||||
|
||||
# initialize latent image guidance
|
||||
cond_fn = None
|
||||
|
||||
preds, stage1_preds = process(
|
||||
preds = process(
|
||||
model, [x], steps=args.steps,
|
||||
t_max=args.t_max, t_min=args.t_min,
|
||||
strength=1,
|
||||
color_fix_type=args.color_fix_type,
|
||||
disable_preprocess_model=args.disable_preprocess_model,
|
||||
cond_fn=cond_fn,
|
||||
tiled=args.tiled, tile_size=args.tile_size, tile_stride=args.tile_stride
|
||||
)
|
||||
pred, stage1_pred = preds[0], stage1_preds[0]
|
||||
|
||||
pred = preds[0]
|
||||
# remove padding
|
||||
pred = pred[:lq_resized.height, :lq_resized.width, :]
|
||||
# pred = pred[:lq_resized.height, :lq_resized.width, :]
|
||||
|
||||
if args.show_lq:
|
||||
pred = np.array(Image.fromarray(pred).resize(lq.size, Image.LANCZOS))
|
||||
stage1_pred = np.array(Image.fromarray(stage1_pred).resize(lq.size, Image.LANCZOS))
|
||||
lq = np.array(lq)
|
||||
images = [lq, pred] if args.disable_preprocess_model else [lq, stage1_pred, pred]
|
||||
images = [lq, pred]
|
||||
Image.fromarray(np.concatenate(images, axis=1)).save(save_path)
|
||||
else:
|
||||
Image.fromarray(pred).resize(lq.size, Image.LANCZOS).save(save_path)
|
||||
# pred.save(save_path)
|
||||
print(f"save to {save_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
Reference in New Issue
Block a user