Upload inference.py with huggingface_hub
Browse files- inference.py +51 -45
inference.py
CHANGED
|
@@ -1,14 +1,15 @@
|
|
| 1 |
#!/usr/bin/env python3
|
| 2 |
-
"""OracleZoom - faithful extreme super-resolution (4x -> 16x -> 64x -> 256x).
|
| 3 |
|
| 4 |
Self-contained: this repo + auto-downloaded Stable Diffusion 3-medium and Qwen2.5-VL-3B.
|
| 5 |
Vendored Chain-of-Zoom code lives in ./coz, checkpoints in ./ckpt, merged model = merged_transformer.safetensors.
|
| 6 |
|
| 7 |
-
Usage:
|
| 8 |
-
python inference.py --input .
|
| 9 |
-
|
|
|
|
| 10 |
"""
|
| 11 |
-
import argparse,
|
| 12 |
import torch
|
| 13 |
from PIL import Image
|
| 14 |
from torchvision import transforms
|
|
@@ -41,8 +42,9 @@ class _SRArgs:
|
|
| 41 |
self.mixed_precision = "fp16"
|
| 42 |
|
| 43 |
|
| 44 |
-
def build_sr(ckpt, sd3, process_size):
|
| 45 |
from osediff_sd3 import OSEDiff_SD3_TEST, SD3Euler
|
|
|
|
| 46 |
sr = SD3Euler()
|
| 47 |
for m in [sr.text_enc_1, sr.text_enc_2, sr.text_enc_3, sr.transformer, sr.vae]:
|
| 48 |
m.to("cuda:0")
|
|
@@ -50,7 +52,14 @@ def build_sr(ckpt, sd3, process_size):
|
|
| 50 |
sr.vae.to("cuda:0", dtype=torch.float32)
|
| 51 |
for m in [sr.text_enc_1, sr.text_enc_2, sr.text_enc_3, sr.transformer, sr.vae]:
|
| 52 |
m.requires_grad_(False)
|
| 53 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 54 |
|
| 55 |
|
| 56 |
def build_vlm(ckpt):
|
|
@@ -77,10 +86,38 @@ def vlm_prompt(model, proc, pvi, first, second, max_new_tokens=32):
|
|
| 77 |
return proc.batch_decode(trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
|
| 78 |
|
| 79 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 80 |
def main():
|
| 81 |
-
ap = argparse.ArgumentParser()
|
| 82 |
-
ap.add_argument("--input", required=True, help="
|
| 83 |
-
ap.add_argument("--output",
|
| 84 |
ap.add_argument("--merged", default=os.path.join(HERE, "merged_transformer.safetensors"))
|
| 85 |
ap.add_argument("--ckpt", default=os.path.join(HERE, "ckpt"))
|
| 86 |
ap.add_argument("--sd3", default="stabilityai/stable-diffusion-3-medium-diffusers")
|
|
@@ -89,43 +126,12 @@ def main():
|
|
| 89 |
ap.add_argument("--process_size", type=int, default=512)
|
| 90 |
ap.add_argument("--max_new_tokens", type=int, default=32)
|
| 91 |
a = ap.parse_args()
|
| 92 |
-
os.makedirs(a.output, exist_ok=True)
|
| 93 |
|
| 94 |
-
sr = build_sr(a.ckpt, a.sd3, a.process_size)
|
| 95 |
-
from safetensors.torch import load_file
|
| 96 |
-
sd = load_file(a.merged)
|
| 97 |
-
dev = next(sr.model.transformer.parameters()).device
|
| 98 |
-
sd = {k: v.to(dev, dtype=torch.float32) for k, v in sd.items()}
|
| 99 |
-
miss, unexp = sr.model.transformer.load_state_dict(sd, strict=False)
|
| 100 |
-
print(f"[OracleZoom] merged transformer loaded (missing={len(miss)} unexpected={len(unexp)})", flush=True)
|
| 101 |
model, proc, pvi = build_vlm(a.ckpt)
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
bname = os.path.splitext(os.path.basename(img_path))[0]
|
| 106 |
-
rec_dir = os.path.join(a.output, "per-sample", bname)
|
| 107 |
-
os.makedirs(rec_dir, exist_ok=True)
|
| 108 |
-
cur = resize_and_center_crop(Image.open(img_path).convert("RGB"), a.process_size)
|
| 109 |
-
cur.save(f"{rec_dir}/0.png")
|
| 110 |
-
for rec in range(a.rec_num):
|
| 111 |
-
prev = Image.open(f"{rec_dir}/{rec}.png").convert("RGB")
|
| 112 |
-
w, h = prev.size
|
| 113 |
-
nw, nh = w // a.upscale, h // a.upscale
|
| 114 |
-
crop = prev.crop(((w - nw) // 2, (h - nh) // 2, (w + nw) // 2, (h + nh) // 2))
|
| 115 |
-
zoom = crop.resize((w, h), Image.BICUBIC)
|
| 116 |
-
zp = f"{rec_dir}/{rec + 1}_input.png"
|
| 117 |
-
zoom.save(zp)
|
| 118 |
-
prompt = vlm_prompt(model, proc, pvi, f"{rec_dir}/{rec}.png", zp, a.max_new_tokens)
|
| 119 |
-
lq = _to_tensor(zoom).unsqueeze(0).to("cuda") * 2 - 1
|
| 120 |
-
with torch.no_grad():
|
| 121 |
-
out = torch.clamp(sr(lq, prompt=prompt)[0].cpu(), -1.0, 1.0)
|
| 122 |
-
transforms.ToPILImage()(out * 0.5 + 0.5).save(f"{rec_dir}/{rec + 1}.png")
|
| 123 |
-
print(f" {bname} scale{rec + 1} ({4 ** (rec + 1)}x): {prompt}", flush=True)
|
| 124 |
-
for s in range(a.rec_num + 1):
|
| 125 |
-
d = os.path.join(a.output, "per-scale", f"scale{s}")
|
| 126 |
-
os.makedirs(d, exist_ok=True)
|
| 127 |
-
Image.open(f"{rec_dir}/{s}.png").save(os.path.join(d, f"{bname}.png"))
|
| 128 |
-
print("[OracleZoom] done ->", a.output, flush=True)
|
| 129 |
|
| 130 |
|
| 131 |
if __name__ == "__main__":
|
|
|
|
| 1 |
#!/usr/bin/env python3
|
| 2 |
+
"""OracleZoom - faithful extreme super-resolution of ONE image (4x -> 16x -> 64x -> 256x).
|
| 3 |
|
| 4 |
Self-contained: this repo + auto-downloaded Stable Diffusion 3-medium and Qwen2.5-VL-3B.
|
| 5 |
Vendored Chain-of-Zoom code lives in ./coz, checkpoints in ./ckpt, merged model = merged_transformer.safetensors.
|
| 6 |
|
| 7 |
+
Usage (one image in, all scales out):
|
| 8 |
+
python inference.py --input photo.jpg --output ./outputs
|
| 9 |
+
Writes: outputs/<name>_1x.png (input crop), _4x.png, _16x.png, _64x.png, _256x.png
|
| 10 |
+
(To batch many images, just call zoom_image() in a loop.)
|
| 11 |
"""
|
| 12 |
+
import argparse, os, sys, tempfile
|
| 13 |
import torch
|
| 14 |
from PIL import Image
|
| 15 |
from torchvision import transforms
|
|
|
|
| 42 |
self.mixed_precision = "fp16"
|
| 43 |
|
| 44 |
|
| 45 |
+
def build_sr(ckpt, sd3, merged, process_size):
|
| 46 |
from osediff_sd3 import OSEDiff_SD3_TEST, SD3Euler
|
| 47 |
+
from safetensors.torch import load_file
|
| 48 |
sr = SD3Euler()
|
| 49 |
for m in [sr.text_enc_1, sr.text_enc_2, sr.text_enc_3, sr.transformer, sr.vae]:
|
| 50 |
m.to("cuda:0")
|
|
|
|
| 52 |
sr.vae.to("cuda:0", dtype=torch.float32)
|
| 53 |
for m in [sr.text_enc_1, sr.text_enc_2, sr.text_enc_3, sr.transformer, sr.vae]:
|
| 54 |
m.requires_grad_(False)
|
| 55 |
+
sr_test = OSEDiff_SD3_TEST(_SRArgs(ckpt, sd3, process_size), sr)
|
| 56 |
+
# load the merged OracleZoom transformer (replaces all transformer weights)
|
| 57 |
+
sd = load_file(merged)
|
| 58 |
+
dev = next(sr_test.model.transformer.parameters()).device
|
| 59 |
+
sd = {k: v.to(dev, dtype=torch.float32) for k, v in sd.items()}
|
| 60 |
+
miss, unexp = sr_test.model.transformer.load_state_dict(sd, strict=False)
|
| 61 |
+
print(f"[OracleZoom] merged transformer loaded (missing={len(miss)} unexpected={len(unexp)})", flush=True)
|
| 62 |
+
return sr_test
|
| 63 |
|
| 64 |
|
| 65 |
def build_vlm(ckpt):
|
|
|
|
| 86 |
return proc.batch_decode(trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
|
| 87 |
|
| 88 |
|
| 89 |
+
def zoom_image(sr, model, proc, pvi, image_path, out_dir,
|
| 90 |
+
rec_num=4, upscale=4, process_size=512, max_new_tokens=32):
|
| 91 |
+
"""Super-resolve ONE image through the recursion; save every scale to out_dir. Returns list of paths."""
|
| 92 |
+
os.makedirs(out_dir, exist_ok=True)
|
| 93 |
+
stem = os.path.splitext(os.path.basename(image_path))[0]
|
| 94 |
+
work = tempfile.mkdtemp()
|
| 95 |
+
resize_and_center_crop(Image.open(image_path).convert("RGB"), process_size).save(f"{work}/0.png")
|
| 96 |
+
for rec in range(rec_num):
|
| 97 |
+
prev = Image.open(f"{work}/{rec}.png").convert("RGB")
|
| 98 |
+
w, h = prev.size
|
| 99 |
+
nw, nh = w // upscale, h // upscale
|
| 100 |
+
crop = prev.crop(((w - nw) // 2, (h - nh) // 2, (w + nw) // 2, (h + nh) // 2))
|
| 101 |
+
zoom = crop.resize((w, h), Image.BICUBIC)
|
| 102 |
+
zoom.save(f"{work}/{rec + 1}_input.png")
|
| 103 |
+
prompt = vlm_prompt(model, proc, pvi, f"{work}/{rec}.png", f"{work}/{rec + 1}_input.png", max_new_tokens)
|
| 104 |
+
lq = _to_tensor(zoom).unsqueeze(0).to("cuda") * 2 - 1
|
| 105 |
+
with torch.no_grad():
|
| 106 |
+
out = torch.clamp(sr(lq, prompt=prompt)[0].cpu(), -1.0, 1.0)
|
| 107 |
+
transforms.ToPILImage()(out * 0.5 + 0.5).save(f"{work}/{rec + 1}.png")
|
| 108 |
+
print(f" scale{rec + 1} ({4 ** (rec + 1)}x): {prompt}", flush=True)
|
| 109 |
+
saved = []
|
| 110 |
+
for s in range(rec_num + 1):
|
| 111 |
+
dst = os.path.join(out_dir, f"{stem}_{4 ** s}x.png")
|
| 112 |
+
Image.open(f"{work}/{s}.png").save(dst)
|
| 113 |
+
saved.append(dst)
|
| 114 |
+
return saved
|
| 115 |
+
|
| 116 |
+
|
| 117 |
def main():
|
| 118 |
+
ap = argparse.ArgumentParser(description="OracleZoom: super-resolve one image to 4x/16x/64x/256x.")
|
| 119 |
+
ap.add_argument("--input", required=True, help="path to ONE input image")
|
| 120 |
+
ap.add_argument("--output", default="./outputs", help="output folder (all scales saved here)")
|
| 121 |
ap.add_argument("--merged", default=os.path.join(HERE, "merged_transformer.safetensors"))
|
| 122 |
ap.add_argument("--ckpt", default=os.path.join(HERE, "ckpt"))
|
| 123 |
ap.add_argument("--sd3", default="stabilityai/stable-diffusion-3-medium-diffusers")
|
|
|
|
| 126 |
ap.add_argument("--process_size", type=int, default=512)
|
| 127 |
ap.add_argument("--max_new_tokens", type=int, default=32)
|
| 128 |
a = ap.parse_args()
|
|
|
|
| 129 |
|
| 130 |
+
sr = build_sr(a.ckpt, a.sd3, a.merged, a.process_size)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 131 |
model, proc, pvi = build_vlm(a.ckpt)
|
| 132 |
+
saved = zoom_image(sr, model, proc, pvi, a.input, a.output,
|
| 133 |
+
a.rec_num, a.upscale, a.process_size, a.max_new_tokens)
|
| 134 |
+
print("[OracleZoom] saved:", *saved, sep="\n ", flush=True)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 135 |
|
| 136 |
|
| 137 |
if __name__ == "__main__":
|