Download inference.py from dipta007/OracleZoom: direct link, hf CLI and curl.
- Browser
- Download file 6.6 kB
-
https://huggingface.co/dipta007/OracleZoom/resolve/refs%2Fpr%2F1/inference.py
- Command line
-
hf download hf://dipta007/OracleZoom@refs/pr/1/inference.py
-
curl -L -o inference.py https://huggingface.co/dipta007/OracleZoom/resolve/refs%2Fpr%2F1/inference.py
6.6 kB
| #!/usr/bin/env python3 | |
| """OracleZoom - faithful extreme super-resolution of ONE image (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 (one image in, all scales out): | |
| python inference.py --input photo.jpg --output ./outputs | |
| Writes: outputs/<name>_1x.png (input crop), _4x.png, _16x.png, _64x.png, _256x.png | |
| (To batch many images, just call zoom_image() in a loop.) | |
| """ | |
| import argparse, os, sys, tempfile | |
| 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, merged, process_size): | |
| from osediff_sd3 import OSEDiff_SD3_TEST, SD3Euler | |
| from safetensors.torch import load_file | |
| 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) | |
| sr_test = OSEDiff_SD3_TEST(_SRArgs(ckpt, sd3, process_size), sr) | |
| # load the merged OracleZoom transformer (replaces all transformer weights) | |
| sd = load_file(merged) | |
| dev = next(sr_test.model.transformer.parameters()).device | |
| sd = {k: v.to(dev, dtype=torch.float32) for k, v in sd.items()} | |
| miss, unexp = sr_test.model.transformer.load_state_dict(sd, strict=False) | |
| print(f"[OracleZoom] merged transformer loaded (missing={len(miss)} unexpected={len(unexp)})", flush=True) | |
| return sr_test | |
| 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 zoom_image(sr, model, proc, pvi, image_path, out_dir, | |
| rec_num=4, upscale=4, process_size=512, max_new_tokens=32): | |
| """Super-resolve ONE image through the recursion; save every scale to out_dir. Returns list of paths.""" | |
| os.makedirs(out_dir, exist_ok=True) | |
| stem = os.path.splitext(os.path.basename(image_path))[0] | |
| work = tempfile.mkdtemp() | |
| resize_and_center_crop(Image.open(image_path).convert("RGB"), process_size).save(f"{work}/0.png") | |
| for rec in range(rec_num): | |
| prev = Image.open(f"{work}/{rec}.png").convert("RGB") | |
| w, h = prev.size | |
| nw, nh = w // upscale, h // upscale | |
| crop = prev.crop(((w - nw) // 2, (h - nh) // 2, (w + nw) // 2, (h + nh) // 2)) | |
| zoom = crop.resize((w, h), Image.BICUBIC) | |
| zoom.save(f"{work}/{rec + 1}_input.png") | |
| prompt = vlm_prompt(model, proc, pvi, f"{work}/{rec}.png", f"{work}/{rec + 1}_input.png", 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"{work}/{rec + 1}.png") | |
| print(f" scale{rec + 1} ({4 ** (rec + 1)}x): {prompt}", flush=True) | |
| saved = [] | |
| for s in range(rec_num + 1): | |
| dst = os.path.join(out_dir, f"{stem}_{4 ** s}x.png") | |
| Image.open(f"{work}/{s}.png").save(dst) | |
| saved.append(dst) | |
| return saved | |
| def main(): | |
| ap = argparse.ArgumentParser(description="OracleZoom: super-resolve one image to 4x/16x/64x/256x.") | |
| ap.add_argument("--input", required=True, help="path to ONE input image") | |
| ap.add_argument("--output", default="./outputs", help="output folder (all scales saved here)") | |
| 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() | |
| sr = build_sr(a.ckpt, a.sd3, a.merged, a.process_size) | |
| model, proc, pvi = build_vlm(a.ckpt) | |
| saved = zoom_image(sr, model, proc, pvi, a.input, a.output, | |
| a.rec_num, a.upscale, a.process_size, a.max_new_tokens) | |
| print("[OracleZoom] saved:", *saved, sep="\n ", flush=True) | |
| if __name__ == "__main__": | |
| main() | |