Download src/inference_4k.py from XenderYang/CSIGv3_train_script: direct link, hf CLI and curl.
- Browser
- Download file 5.49 kB
-
https://huggingface.co/XenderYang/CSIGv3_train_script/resolve/main/src/inference_4k.py
- Command line
-
hf download hf://XenderYang/CSIGv3_train_script/src/inference_4k.py
-
curl -L -o inference_4k.py https://huggingface.co/XenderYang/CSIGv3_train_script/resolve/main/src/inference_4k.py
5.49 kB
| #!/usr/bin/env python | |
| """4K 整图增强: 512 窗 + overlap 128 线性融合 + AdaIN 对齐 LR(禁止 zero-padding, 边缘 reflect)。 | |
| x1 语义: 对每个 512 窗先缩到 128, 走官方 4x 学生模型, 回 512。 | |
| 用法: python src/inference_4k.py --lr_dir <dir> --net weight/s2/net_params_X.pkl --out output_dir | |
| """ | |
| import argparse, copy, os, sys | |
| from pathlib import Path | |
| REPO = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(REPO)); sys.path.insert(0, str(REPO / "src")); sys.path.insert(0, str(REPO / "official")) | |
| import torch, torch.nn.functional as F | |
| import numpy as np | |
| from PIL import Image | |
| from common import load_diffusers_sd, load_pruned_decoder, build_net | |
| def build_model(net_pkl, half_decoder, model_id, device, bf16=True): | |
| dtype = torch.bfloat16 if (bf16 and device == "cuda") else torch.float32 | |
| vae, unet, _, _ = load_diffusers_sd(model_id, dtype=torch.float32, device="cpu") | |
| del vae | |
| decoder = load_pruned_decoder(half_decoder, device="cpu", dtype=torch.float32) | |
| net = build_net(unet, decoder) | |
| net.to(device=device, dtype=dtype) | |
| sd = torch.load(net_pkl, map_location="cpu", weights_only=False) | |
| if any(k.startswith("module.") for k in sd): | |
| sd = {k.replace("module.", "", 1): v for k, v in sd.items()} | |
| net.load_state_dict(sd, strict=True) | |
| net.eval() | |
| tail = torch.nn.Sequential(*decoder.up_blocks, decoder.conv_norm_out, | |
| decoder.conv_act, decoder.conv_out).to(device=device, dtype=dtype).eval() | |
| return net, tail | |
| def infer_tile(net, tail, lr128, lr_stats): | |
| """lr128: [1,3,128,128] [-1,1]; 返回 [1,3,512,512] [-1,1] 并 AdaIN 对齐 lr_stats""" | |
| with torch.no_grad(): | |
| z = net(lr128) | |
| out = tail(z) | |
| out = out.float() | |
| mu, std = out.mean(dim=(2, 3), keepdim=True), out.std(dim=(2, 3), keepdim=True) | |
| lmu, lstd = lr_stats | |
| out = (out - mu) / (std + 1e-6) * lstd + lmu | |
| return out.clamp(-1, 1) | |
| def infer_4k_single(lr_path, net, tail, device, tile=512, overlap=128, bf16=True): | |
| im = Image.open(lr_path).convert("RGB") | |
| w, h = im.size | |
| a = np.asarray(im, dtype=np.float32) / 255.0 # [H,W,3] | |
| # reflect pad so tile 整除 | |
| stride = tile - overlap | |
| pad_r = (stride - w % stride) % stride + (tile - stride) if (w % stride) else max(0, tile - stride) | |
| pad_b = (stride - h % stride) % stride + (tile - stride) if (h % stride) else max(0, tile - stride) | |
| a = np.pad(a, ((0, pad_b), (0, pad_r), (0, 0)), mode="reflect") | |
| H, W = a.shape[:2] | |
| acc = np.zeros((H, W, 3), dtype=np.float64) | |
| wsum = np.zeros((H, W, 1), dtype=np.float64) | |
| rows = list(range(0, H - tile + 1, stride)) or [0] | |
| cols = list(range(0, W - tile + 1, stride)) or [0] | |
| if rows[-1] + tile < H: rows.append(H - tile) | |
| if cols[-1] + tile < W: cols.append(W - tile) | |
| # 线性窗权重 | |
| ramp = np.minimum(np.arange(tile), np.arange(tile)[::-1]) / (tile // 2) | |
| w2d = (ramp[None, :] * ramp[:, None])[..., None].astype(np.float64) | |
| for y in rows: | |
| for x in cols: | |
| crop = a[y:y + tile, x:x + tile] | |
| lr128 = np.asarray(Image.fromarray((crop * 255).astype(np.uint8)).resize((128, 128), Image.BICUBIC), | |
| dtype=np.float32) / 255.0 | |
| lr_t = torch.from_numpy(lr128.transpose(2, 0, 1))[None].to(device) * 2 - 1 | |
| stats = (torch.tensor(crop.mean(axis=(0, 1))[None, :, None, None], device=device) * 2 - 1, | |
| torch.tensor(crop.std(axis=(0, 1))[None, :, None, None], device=device)) | |
| with torch.autocast("cuda", dtype=torch.bfloat16, enabled=bf16): | |
| out = infer_tile(net, tail, lr_t, stats) | |
| o = out[0].float().cpu().numpy().transpose(1, 2, 0) # [512,512,3] [-1,1] | |
| o = (o + 1) / 2.0 | |
| acc[y:y + tile, x:x + tile] += o.astype(np.float64) * w2d | |
| wsum[y:y + tile, x:x + tile] += w2d | |
| res = acc / (wsum + 1e-9) | |
| res = res[:h, :w] | |
| return Image.fromarray((np.clip(res, 0, 1) * 255).astype(np.uint8)) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--lr_dir", required=True) | |
| ap.add_argument("--net", required=True) | |
| ap.add_argument("--out", required=True) | |
| ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt") | |
| ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base") | |
| ap.add_argument("--tile", type=int, default=512) | |
| ap.add_argument("--overlap", type=int, default=128) | |
| ap.add_argument("--no_bf16", action="store_true") | |
| args = ap.parse_args() | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| os.makedirs(args.out, exist_ok=True) | |
| net, tail = build_model(args.net, args.half_decoder, args.model_id, device, bf16=not args.no_bf16) | |
| files = sorted([f for f in os.listdir(args.lr_dir) if f.lower().endswith((".jpg", ".jpeg", ".png"))]) | |
| import time | |
| t0 = time.time() | |
| for i, f in enumerate(files): | |
| out_im = infer_4k_single(os.path.join(args.lr_dir, f), net, tail, device, | |
| tile=args.tile, overlap=args.overlap, bf16=not args.no_bf16) | |
| out_im.save(os.path.join(args.out, f), quality=95) | |
| if (i + 1) % 10 == 0: | |
| print(f" {i+1}/{len(files)} elapsed {time.time()-t0:.1f}s", flush=True) | |
| print(f"done {len(files)} imgs in {time.time()-t0:.1f}s") | |
| if __name__ == "__main__": | |
| main() | |