#!/usr/bin/env python3 """OracleZoom - faithful extreme super-resolution (4x -> 16x -> 64x -> 256x). Self-contained: this repo + auto-downloaded Stable Diffusion 3-medium and Qwen2.5-VL-3B. Vendored Chain-of-Zoom code lives in ./coz, checkpoints in ./ckpt, merged model = merged_transformer.safetensors. Usage: python inference.py --input ./inputs --output ./outputs Outputs: outputs/per-scale/scale1..4/.png (4x / 16x / 64x / 256x) """ import argparse, glob, os, sys import torch from PIL import Image from torchvision import transforms HERE = os.path.dirname(os.path.abspath(__file__)) sys.path.insert(0, os.path.join(HERE, "coz")) # vendored Chain-of-Zoom modules COZ_PROMPT = ("The second image is a zoom-in of the first image. Based on this knowledge, " "what is in the second image? Give me a set of words.") _to_tensor = transforms.Compose([transforms.ToTensor()]) def resize_and_center_crop(img, size): w, h = img.size scale = size / min(w, h) nw, nh = int(w * scale), int(h * scale) img = img.resize((nw, nh), Image.LANCZOS) l, t = (nw - size) // 2, (nh - size) // 2 return img.crop((l, t, l + size, t + size)) class _SRArgs: def __init__(self, ckpt, sd3, process_size): self.lora_path = f"{ckpt}/SR_LoRA/model_20001.pkl" self.vae_path = f"{ckpt}/SR_VAE/vae_encoder_20001.pt" self.pretrained_model_name_or_path = sd3 self.process_size = process_size self.lora_rank = 4 self.merge_and_unload_lora = False self.mixed_precision = "fp16" def build_sr(ckpt, sd3, process_size): from osediff_sd3 import OSEDiff_SD3_TEST, SD3Euler sr = SD3Euler() for m in [sr.text_enc_1, sr.text_enc_2, sr.text_enc_3, sr.transformer, sr.vae]: m.to("cuda:0") sr.transformer.to("cuda:0", dtype=torch.float32) sr.vae.to("cuda:0", dtype=torch.float32) for m in [sr.text_enc_1, sr.text_enc_2, sr.text_enc_3, sr.transformer, sr.vae]: m.requires_grad_(False) return OSEDiff_SD3_TEST(_SRArgs(ckpt, sd3, process_size), sr) def build_vlm(ckpt): from transformers import Qwen2_5_VLForConditionalGeneration, AutoProcessor from qwen_vl_utils import process_vision_info from peft import PeftModel name = "Qwen/Qwen2.5-VL-3B-Instruct" model = Qwen2_5_VLForConditionalGeneration.from_pretrained( name, torch_dtype="auto", device_map="auto", attn_implementation="sdpa") proc = AutoProcessor.from_pretrained(name) model = PeftModel.from_pretrained(model, f"{ckpt}/VLM_LoRA/checkpoint-10000").merge_and_unload().eval() return model, proc, process_vision_info def vlm_prompt(model, proc, pvi, first, second, max_new_tokens=32): messages = [{"role": "system", "content": COZ_PROMPT}, {"role": "user", "content": [{"type": "image", "image": first}, {"type": "image", "image": second}]}] text = proc.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) ii, vi = pvi(messages) inputs = proc(text=[text], images=ii, videos=vi, padding=True, return_tensors="pt").to("cuda") gen = model.generate(**inputs, max_new_tokens=max_new_tokens) trimmed = [o[len(i):] for i, o in zip(inputs.input_ids, gen)] return proc.batch_decode(trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0] def main(): ap = argparse.ArgumentParser() ap.add_argument("--input", required=True, help="folder of input images") ap.add_argument("--output", required=True, help="output folder") ap.add_argument("--merged", default=os.path.join(HERE, "merged_transformer.safetensors")) ap.add_argument("--ckpt", default=os.path.join(HERE, "ckpt")) ap.add_argument("--sd3", default="stabilityai/stable-diffusion-3-medium-diffusers") ap.add_argument("--rec_num", type=int, default=4) ap.add_argument("--upscale", type=int, default=4) ap.add_argument("--process_size", type=int, default=512) ap.add_argument("--max_new_tokens", type=int, default=32) a = ap.parse_args() os.makedirs(a.output, exist_ok=True) sr = build_sr(a.ckpt, a.sd3, a.process_size) from safetensors.torch import load_file sd = load_file(a.merged) dev = next(sr.model.transformer.parameters()).device sd = {k: v.to(dev, dtype=torch.float32) for k, v in sd.items()} miss, unexp = sr.model.transformer.load_state_dict(sd, strict=False) print(f"[OracleZoom] merged transformer loaded (missing={len(miss)} unexpected={len(unexp)})", flush=True) model, proc, pvi = build_vlm(a.ckpt) imgs = sorted(p for e in ("*.png", "*.jpg", "*.jpeg", "*.webp") for p in glob.glob(f"{a.input}/{e}")) for img_path in imgs: bname = os.path.splitext(os.path.basename(img_path))[0] rec_dir = os.path.join(a.output, "per-sample", bname) os.makedirs(rec_dir, exist_ok=True) cur = resize_and_center_crop(Image.open(img_path).convert("RGB"), a.process_size) cur.save(f"{rec_dir}/0.png") for rec in range(a.rec_num): prev = Image.open(f"{rec_dir}/{rec}.png").convert("RGB") w, h = prev.size nw, nh = w // a.upscale, h // a.upscale crop = prev.crop(((w - nw) // 2, (h - nh) // 2, (w + nw) // 2, (h + nh) // 2)) zoom = crop.resize((w, h), Image.BICUBIC) zp = f"{rec_dir}/{rec + 1}_input.png" zoom.save(zp) prompt = vlm_prompt(model, proc, pvi, f"{rec_dir}/{rec}.png", zp, a.max_new_tokens) lq = _to_tensor(zoom).unsqueeze(0).to("cuda") * 2 - 1 with torch.no_grad(): out = torch.clamp(sr(lq, prompt=prompt)[0].cpu(), -1.0, 1.0) transforms.ToPILImage()(out * 0.5 + 0.5).save(f"{rec_dir}/{rec + 1}.png") print(f" {bname} scale{rec + 1} ({4 ** (rec + 1)}x): {prompt}", flush=True) for s in range(a.rec_num + 1): d = os.path.join(a.output, "per-scale", f"scale{s}") os.makedirs(d, exist_ok=True) Image.open(f"{rec_dir}/{s}.png").save(os.path.join(d, f"{bname}.png")) print("[OracleZoom] done ->", a.output, flush=True) if __name__ == "__main__": main()