Download inference.py from sharky172/manga-light-colorizer: direct link, hf CLI and curl.
- Browser
- Download file 13.4 kB
-
https://huggingface.co/sharky172/manga-light-colorizer/resolve/main/inference.py
- Command line
-
hf download hf://sharky172/manga-light-colorizer/inference.py
-
curl -L -o inference.py https://huggingface.co/sharky172/manga-light-colorizer/resolve/main/inference.py
13.4 kB
| #!/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() | |