#!/usr/bin/env python3 """ Manga Light Colorizer - Gradio app (local ONNX inference) Runs entirely inside a Hugging Face Space (or locally) using onnxruntime. No external API is called: the v6 generator + SAM encoder ONNX models are loaded from the local `models/` folder and run directly. Models (auto-detected, relative to this script): standalone/models/v6_generator.onnx standalone/models/v6_sam_encoder.onnx Launch: python app.py """ import sys import time from pathlib import Path import cv2 import gradio as gr import numpy as np from PIL import Image try: import onnxruntime as ort except ImportError: print("Error: onnxruntime not installed. Install with: pip install onnxruntime") sys.exit(1) print(f"[startup] Python {sys.version}", flush=True) print(f"[startup] gradio version: {gr.__version__}", flush=True) print(f"[startup] onnxruntime version: {ort.__version__}", flush=True) # ============================================================================ # CONFIG # ============================================================================ SCRIPT_DIR = Path(__file__).resolve().parent GENERATOR_PATH = SCRIPT_DIR / "models" / "v6_generator.onnx" SAM_PATH = SCRIPT_DIR / "models" / "v6_sam_encoder.onnx" EXAMPLES_DIR = SCRIPT_DIR / "input" INFER_SIZE_OPTIONS = [512, 768, 1024] DEFAULT_INFER_SIZE = 768 # ============================================================================ # CORE ONNX INFERENCE (ported from inference.py) # ============================================================================ def denormalize_rgb(rgb_norm: np.ndarray) -> np.ndarray: """[-1, 1] -> [0, 255] uint8.""" return np.clip((rgb_norm + 1.0) * 127.5, 0, 255).astype(np.uint8) def extract_sam_features_onnx(sam_session: ort.InferenceSession, L_bw_norm: np.ndarray): """ Extract SAM features via ONNX. WD14 is intentionally DISABLED (zeros). Args: sam_session: ONNX Runtime session for SAM encoder L_bw_norm: (H, W) grayscale in [-1, 1] Returns: sam_level0, sam_level1, wd14_embedding (all numpy) """ L_01 = (L_bw_norm + 1.0) / 2.0 # [-1,1] -> [0,1] L_1024 = cv2.resize(L_01, (1024, 1024), interpolation=cv2.INTER_LINEAR) rgb_sam = np.stack([L_1024, L_1024, L_1024], axis=0)[np.newaxis].astype(np.float32) sam_out = sam_session.run(None, {"rgb_input": rgb_sam}) sam_level0 = sam_out[0] # (1, 256, 64, 64) sam_level1 = sam_out[1] # (1, 256, 32, 32) wd14_embedding = np.zeros((1, 1024), dtype=np.float32) return sam_level0, sam_level1, wd14_embedding def colorize_onnx( session: ort.InferenceSession, L_bw: np.ndarray, sam_level0: np.ndarray, sam_level1: np.ndarray, wd14_embedding: np.ndarray, ) -> np.ndarray: """Run generator ONNX inference. Returns RGB (H, W, 3) in [0, 255].""" L_norm = (L_bw.astype(np.float32) / 127.5) - 1.0 L_tensor = L_norm[np.newaxis, np.newaxis, :, :] # (1, 1, H, W) ort_inputs = { "L_bw": L_tensor, "sam_level0": sam_level0, "sam_level1": sam_level1, "wd14_embedding": wd14_embedding, } rgb_pred = session.run(None, ort_inputs)[0] # (1, 3, H, W) rgb_pred = rgb_pred[0].transpose(1, 2, 0) # (H, W, 3) return denormalize_rgb(rgb_pred) # ============================================================================ # MODEL LOADING (once, at startup) # ============================================================================ def load_sessions(): """Load generator (+ optional SAM) ONNX sessions. Prefers CUDA if available.""" available = ort.get_available_providers() if "CUDAExecutionProvider" in available: providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] else: providers = ["CPUExecutionProvider"] if not GENERATOR_PATH.exists(): raise FileNotFoundError(f"Generator ONNX not found: {GENERATOR_PATH}") print(f"[startup] Loading generator: {GENERATOR_PATH}", flush=True) session = ort.InferenceSession(str(GENERATOR_PATH), providers=providers) print(f"[startup] Generator provider: {session.get_providers()[0]}", flush=True) sam_session = None if SAM_PATH.exists(): print(f"[startup] Loading SAM encoder: {SAM_PATH}", flush=True) sam_session = ort.InferenceSession(str(SAM_PATH), providers=providers) print("[startup] SAM encoder loaded", flush=True) else: print("[startup] SAM encoder NOT found -> using zeros", flush=True) return session, sam_session SESSION, SAM_SESSION = load_sessions() HAS_SAM = SAM_SESSION is not None # ============================================================================ # GRADIO INFERENCE HANDLER # ============================================================================ def colorize_image(input_image: Image.Image, infer_size: int): """ Colorize a grayscale manga image using local ONNX models. Args: input_image: PIL Image (any mode). infer_size: Square inference resolution. Returns: (colorized PIL Image or None, status message). """ if input_image is None: return None, "⚠️ Please upload an image first." t_start = time.time() # PIL -> grayscale numpy gray = np.array(input_image.convert("L")) orig_H, orig_W = gray.shape infer_size = int(infer_size) # Always resize input to infer_size for inference L_bw = cv2.resize(gray, (infer_size, infer_size), interpolation=cv2.INTER_AREA) H_in, W_in = L_bw.shape L_norm = (L_bw.astype(np.float32) / 127.5) - 1.0 if HAS_SAM: sam_level0, sam_level1, wd14_embedding = extract_sam_features_onnx(SAM_SESSION, L_norm) else: sam_level0 = np.zeros((1, 256, H_in // 16, W_in // 16), dtype=np.float32) sam_level1 = np.zeros((1, 256, H_in // 32, W_in // 32), dtype=np.float32) wd14_embedding = np.zeros((1, 1024), dtype=np.float32) rgb_output = colorize_onnx(SESSION, L_bw, sam_level0, sam_level1, wd14_embedding) # (infer_size, infer_size, 3) -> back to original input resolution rgb_output = cv2.resize(rgb_output, (orig_W, orig_H), interpolation=cv2.INTER_LANCZOS4) result = Image.fromarray(rgb_output) elapsed = time.time() - t_start status = ( f"✅ Colorization complete! " f"({orig_W}×{orig_H} px, infer {infer_size}×{infer_size}, {elapsed:.2f}s)" ) return result, status # ============================================================================ # GRADIO UI # ============================================================================ def collect_examples(): """Build example list from the input/ folder.""" examples = [] if EXAMPLES_DIR.is_dir(): for ext in ("*.jpg", "*.jpeg", "*.png", "*.bmp", "*.webp"): for f in sorted(EXAMPLES_DIR.glob(ext)): examples.append([str(f), DEFAULT_INFER_SIZE]) return examples def build_interface() -> gr.Blocks: with gr.Blocks( title="Manga Light Colorizer", theme=gr.themes.Soft(), ) as demo: gr.Markdown( """ # 🎨 Manga Light Colorizer Upload a black-and-white manga image and let the AI bring it to life in color. > Runs **fully locally** with ONNX Runtime — no external API call. > The model was trained at **512×512**; the further the inference resolution > differs from 512, the less faithful the colors may be. """ ) with gr.Row(): with gr.Column(scale=1): input_image = gr.Image(label="Input Image", type="pil") infer_size = gr.Radio( choices=INFER_SIZE_OPTIONS, value=DEFAULT_INFER_SIZE, label="Inference Resolution", info=( "Square resolution used for inference. Output is resized back " "to the original input resolution. 512 = best color fidelity." ), ) colorize_btn = gr.Button("🎨 Colorize", variant="primary", size="lg") with gr.Column(scale=1): output_image = gr.Image( label="Colorized Output", type="pil", interactive=False, ) status_text = gr.Textbox(label="Status", interactive=False, lines=2) colorize_btn.click( fn=colorize_image, inputs=[input_image, infer_size], outputs=[output_image, status_text], ) examples = collect_examples() if examples: gr.Examples( examples=examples, inputs=[input_image, infer_size], outputs=[output_image, status_text], fn=colorize_image, cache_examples=False, ) gr.Markdown( """ --- ### 📝 Notes - Supported input formats: **JPEG, PNG, WebP, BMP**. - Inference runs locally via **ONNX Runtime** (CUDA if available, else CPU). - Pipeline: `grayscale → resize → SAM encoder → generator → resize to original`. """ ) return demo print("[startup] calling build_interface()...", flush=True) demo = build_interface() print(f"[startup] demo object created: {demo}", flush=True) if __name__ == "__main__": print("[startup] running as __main__, calling demo.launch()", flush=True) demo.launch()