manga-light-colorizer / inference.py
sharky172's picture
Upload 19 files
a7f6e8c verified
Raw History Blame Contribute Delete
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()