dipta007 commited on
Commit
b0f1186
·
verified ·
1 Parent(s): 9e24d5c

Upload inference.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. 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 ./inputs --output ./outputs
9
- Outputs: outputs/per-scale/scale1..4/<name>.png (4x / 16x / 64x / 256x)
 
10
  """
11
- import argparse, glob, os, sys
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
- return OSEDiff_SD3_TEST(_SRArgs(ckpt, sd3, process_size), sr)
 
 
 
 
 
 
 
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="folder of input images")
83
- ap.add_argument("--output", required=True, help="output folder")
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
- imgs = sorted(p for e in ("*.png", "*.jpg", "*.jpeg", "*.webp") for p in glob.glob(f"{a.input}/{e}"))
104
- for img_path in imgs:
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__":