#!/usr/bin/env python3 """ Manga Light Colorizer - V6 ONNX Inference Script No PyTorch — only ONNX runtime + numpy + cv2. All semantic inputs are zeros (no semantic guidance). Models auto-detected from models/ folder (relative to script location): standalone/models/v6_generator.onnx standalone/models/v6_sam_encoder.onnx Usage: python standalone/inference.py --input input.png python standalone/inference.py --input input.png --infer-size 1024 python standalone/inference.py --input ./input_folder/ python standalone/inference.py --input input.png --output_dir ./output/ """ import argparse import glob import sys import time from pathlib import Path import cv2 import numpy as np try: import onnxruntime as ort except ImportError: print("Error: onnxruntime not installed. Install with: pip install onnxruntime") sys.exit(1) # Supported image extensions IMAGE_EXTENSIONS = {'*.jpg', '*.jpeg', '*.png', '*.bmp', '*.tiff', '*.tif', '*.webp'} # ============================================================================ # STABLE UTILS (minimal, no external dependencies) # ============================================================================ 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 load_image(image_path: str, target_size: int = None, pad_to_multiple: int = 32): """ Load and preprocess grayscale image. If target_size is specified, resize to that square size. Otherwise, pad to nearest multiple of pad_to_multiple (preserves original resolution). Returns: (processed_img, original_size, padding) - original_size: (orig_W, orig_H) - padding: (pad_bottom, pad_right) applied """ img = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE) if img is None: raise ValueError(f"Failed to load image: {image_path}") original_size = (img.shape[1], img.shape[0]) # (W, H) if target_size is not None: # Resize to fixed square size img = cv2.resize(img, (target_size, target_size), interpolation=cv2.INTER_AREA) return img, original_size, (0, 0) # Pad to multiple of pad_to_multiple (required by 5-level encoder) H, W = img.shape pad_h = (pad_to_multiple - H % pad_to_multiple) % pad_to_multiple pad_w = (pad_to_multiple - W % pad_to_multiple) % pad_to_multiple if pad_h > 0 or pad_w > 0: img = np.pad(img, ((0, pad_h), (0, pad_w)), mode='reflect') return img, original_size, (pad_h, pad_w) 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) """ H, W = L_bw_norm.shape # SAM: expects (B, 3, 1024, 1024) RGB in [0, 1] L_01 = (L_bw_norm + 1.0) / 2.0 # [-1,1] -> [0,1] # Resize to 1024x1024 for SAM 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] # (1, 3, 1024, 1024) rgb_sam = rgb_sam.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: DISABLED - return zeros 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 ONNX inference. Args: session: ONNX Runtime session L_bw: (H, W) grayscale in [0, 255] sam_level0: (1, 256, Hs0, Ws0) float32 sam_level1: (1, 256, Hs1, Ws1) float32 wd14_embedding: (1, 1024) float32 (zeros - WD14 disabled) Returns: RGB output (H, W, 3) in [0, 255] """ # Normalize L_bw to [-1, 1] L_norm = (L_bw.astype(np.float32) / 127.5) - 1.0 L_tensor = L_norm[np.newaxis, np.newaxis, :, :] # (1, 1, H, W) # Run ONNX 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) # Convert to (H, W, 3) uint8 rgb_pred = rgb_pred[0].transpose(1, 2, 0) # (H, W, 3) rgb_output = denormalize_rgb(rgb_pred) return rgb_output def get_output_path(input_path: Path, output_folder: Path, input_name: str) -> Path: """Generate output path from input path, preserving filename in output folder.""" return output_folder / input_name def collect_input_files(input_path: str) -> list: """Collect all image files from input (file or folder). Returns list of Path objects.""" input_p = Path(input_path) files = [] if input_p.is_file(): files.append(input_p) elif input_p.is_dir(): for ext in IMAGE_EXTENSIONS: files.extend(Path(f) for f in glob.glob(str(input_p / ext), recursive=True)) # Sort for deterministic ordering files.sort() else: raise ValueError(f"Input not found: {input_path}") return files def process_image( image_path: Path, session: ort.InferenceSession, sam_session: ort.InferenceSession, output_folder: Path, has_sam: bool, infer_size: int = 512, ort_device: str = 'cpu' ) -> tuple: """ Process a single image. Args: infer_size: Resolution to run inference at (square). Input is ALWAYS resized to this. Output is ALWAYS resized back to original input resolution. Returns: (output_path, time_taken, success) """ t_start = time.time() # Load original image (preserve original size before any transformation) img_original = cv2.imread(str(image_path), cv2.IMREAD_GRAYSCALE) if img_original is None: raise ValueError(f"Failed to load image: {image_path}") orig_W, orig_H = img_original.shape[1], img_original.shape[0] # ALWAYS resize input to infer_size for inference L_bw = cv2.resize(img_original, (infer_size, infer_size), interpolation=cv2.INTER_AREA) H_in, W_in = L_bw.shape # Extract features L_norm = (L_bw.astype(np.float32) / 127.5) - 1.0 # [-1, 1] if has_sam and sam_session is not None: 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) # Colorize rgb_output = colorize_onnx(session, L_bw, sam_level0, sam_level1, wd14_embedding) # rgb_output is now (infer_size, infer_size, 3) # ALWAYS resize back to original input resolution rgb_output = cv2.resize(rgb_output, (orig_W, orig_H), interpolation=cv2.INTER_LANCZOS4) # Save output output_path = get_output_path(image_path, output_folder, image_path.name) output_path.parent.mkdir(parents=True, exist_ok=True) cv2.imwrite(str(output_path), cv2.cvtColor(rgb_output, cv2.COLOR_RGB2BGR)) t_end = time.time() return output_path, t_end - t_start, True # ============================================================================ # MAIN # ============================================================================ def main(): parser = argparse.ArgumentParser( description='Manga Light Colorizer - V6 ONNX Inference', formatter_class=argparse.RawDescriptionHelpFormatter, epilog=""" Examples: # Single file python inference.py input.png --onnx-model v6_generator.onnx # Folder (all images) python inference.py ./input_folder/ --onnx-model v6_generator.onnx # With SAM python inference.py input.png --onnx-model v6_generator.onnx --sam-onnx v6_sam_encoder.onnx # Custom output folder python inference.py input.png --onnx-model v6_generator.onnx --output ./output_folder/ """ ) parser.add_argument('--input', type=str, required=True, help='Input grayscale image or folder of images') parser.add_argument('--onnx-model', type=str, default=None, help='Path to ONNX model (default: models/v6_generator.onnx relative to script)') parser.add_argument('--sam-onnx', type=str, default=None, help='SAM ONNX model path (default: models/v6_sam_encoder.onnx relative to script)') parser.add_argument('--output_dir', type=str, default='./output/', help='Output folder (default: ./output/)') parser.add_argument('--infer-size', type=int, default=768, help='Inference resolution (square). Default: 768. Input is resized to this for inference, output is resized back to original.') parser.add_argument('--ort-device', type=str, default='cpu', choices=['cpu', 'cuda'], help='ONNX Runtime device (default: cpu)') args = parser.parse_args() # Auto-detect models from models/ folder (relative to script location) script_dir = Path(__file__).resolve().parent if args.onnx_model is None: args.onnx_model = str(script_dir / 'models' / 'v6_generator.onnx') if args.sam_onnx is None: args.sam_onnx = str(script_dir / 'models' / 'v6_sam_encoder.onnx') # Validate inputs if not Path(args.input).exists(): print(f"Error: Input not found: {args.input}") sys.exit(1) if not Path(args.onnx_model).exists(): print(f"Error: ONNX model not found: {args.onnx_model}") sys.exit(1) # Setup output folder output_folder = Path(args.output_dir) output_folder.mkdir(parents=True, exist_ok=True) # Detect SAM availability has_sam = Path(args.sam_onnx).exists() print("=" * 60) # 1. Detect model paths print(f"[1/4] Detected models:") print(f" Generator: {args.onnx_model}") if has_sam: print(f" SAM: {args.sam_onnx}") else: print(f" SAM: (disabled)") # 2. Load ONNX model print(f"[2/4] Loading generator: {args.onnx_model}") providers = ['CUDAExecutionProvider', 'CPUExecutionProvider'] if args.ort_device == 'cuda' else ['CPUExecutionProvider'] session = ort.InferenceSession(args.onnx_model, providers=providers) active_provider = session.get_providers()[0] print(f" OK Provider: {active_provider}") # 3. Load SAM (if provided) sam_session = None if has_sam: print(f"[3/4] Loading SAM ONNX: {args.sam_onnx}") sam_session = ort.InferenceSession(args.sam_onnx, providers=providers) print(f" OK SAM loaded") else: print(f"[3/4] SAM: DISABLED (using zeros)") # 4. Collect input files print(f"[4/4] Processing images...") try: input_files = collect_input_files(args.input) except ValueError as e: print(f"Error: {e}") sys.exit(1) if not input_files: print(f"Error: No image files found in: {args.input}") sys.exit(1) print(f" Found {len(input_files)} image(s)") print(f" Output folder: {output_folder.resolve()}") print() # Process all images total_time = 0 success_count = 0 fail_count = 0 print(f" Inference size: {args.infer_size}x{args.infer_size} (output resized to original)") print() for i, img_path in enumerate(input_files, 1): print(f"[{i}/{len(input_files)}] Processing: {img_path.name}") try: out_path, elapsed, ok = process_image( img_path, session, sam_session, output_folder, has_sam, infer_size=args.infer_size, ort_device=args.ort_device ) if ok: print(f" OK -> {out_path} ({elapsed:.2f}s)") success_count += 1 else: print(f" FAILED") fail_count += 1 except Exception as e: print(f" Error: {e}") fail_count += 1 total_time += elapsed # Summary print() print(f"{'=' * 60}") print(f"Summary: {success_count}/{len(input_files)} succeeded, {fail_count} failed") print(f"Total time: {total_time:.2f}s") print(f"Output folder: {output_folder.resolve()}") print(f"{'=' * 60}") if __name__ == '__main__': main()