#!/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
--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()