File size: 5,494 Bytes
4811c23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
#!/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()