sharky172 commited on
Commit
2fb022c
Β·
verified Β·
1 Parent(s): 7676098

Upload 4 files

Browse files
Files changed (4) hide show
  1. app.py +277 -0
  2. packages.txt +1 -0
  3. requirements-cuda.txt +4 -0
  4. requirements.txt +5 -0
app.py ADDED
@@ -0,0 +1,277 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Manga Light Colorizer - Gradio app (local ONNX inference)
4
+
5
+ Runs entirely inside a Hugging Face Space (or locally) using onnxruntime.
6
+ No external API is called: the v6 generator + SAM encoder ONNX models are
7
+ loaded from the local `models/` folder and run directly.
8
+
9
+ Models (auto-detected, relative to this script):
10
+ standalone/models/v6_generator.onnx
11
+ standalone/models/v6_sam_encoder.onnx
12
+
13
+ Launch:
14
+ python app.py
15
+ """
16
+
17
+ import sys
18
+ import time
19
+ from pathlib import Path
20
+
21
+ import cv2
22
+ import gradio as gr
23
+ import numpy as np
24
+ from PIL import Image
25
+
26
+ try:
27
+ import onnxruntime as ort
28
+ except ImportError:
29
+ print("Error: onnxruntime not installed. Install with: pip install onnxruntime")
30
+ sys.exit(1)
31
+
32
+ print(f"[startup] Python {sys.version}", flush=True)
33
+ print(f"[startup] gradio version: {gr.__version__}", flush=True)
34
+ print(f"[startup] onnxruntime version: {ort.__version__}", flush=True)
35
+
36
+
37
+ # ============================================================================
38
+ # CONFIG
39
+ # ============================================================================
40
+
41
+ SCRIPT_DIR = Path(__file__).resolve().parent
42
+ GENERATOR_PATH = SCRIPT_DIR / "models" / "v6_generator.onnx"
43
+ SAM_PATH = SCRIPT_DIR / "models" / "v6_sam_encoder.onnx"
44
+ EXAMPLES_DIR = SCRIPT_DIR / "input"
45
+
46
+ INFER_SIZE_OPTIONS = [512, 768, 1024]
47
+ DEFAULT_INFER_SIZE = 768
48
+
49
+
50
+ # ============================================================================
51
+ # CORE ONNX INFERENCE (ported from inference.py)
52
+ # ============================================================================
53
+
54
+ def denormalize_rgb(rgb_norm: np.ndarray) -> np.ndarray:
55
+ """[-1, 1] -> [0, 255] uint8."""
56
+ return np.clip((rgb_norm + 1.0) * 127.5, 0, 255).astype(np.uint8)
57
+
58
+
59
+ def extract_sam_features_onnx(sam_session: ort.InferenceSession, L_bw_norm: np.ndarray):
60
+ """
61
+ Extract SAM features via ONNX. WD14 is intentionally DISABLED (zeros).
62
+
63
+ Args:
64
+ sam_session: ONNX Runtime session for SAM encoder
65
+ L_bw_norm: (H, W) grayscale in [-1, 1]
66
+
67
+ Returns:
68
+ sam_level0, sam_level1, wd14_embedding (all numpy)
69
+ """
70
+ L_01 = (L_bw_norm + 1.0) / 2.0 # [-1,1] -> [0,1]
71
+ L_1024 = cv2.resize(L_01, (1024, 1024), interpolation=cv2.INTER_LINEAR)
72
+ rgb_sam = np.stack([L_1024, L_1024, L_1024], axis=0)[np.newaxis].astype(np.float32)
73
+
74
+ sam_out = sam_session.run(None, {"rgb_input": rgb_sam})
75
+ sam_level0 = sam_out[0] # (1, 256, 64, 64)
76
+ sam_level1 = sam_out[1] # (1, 256, 32, 32)
77
+
78
+ wd14_embedding = np.zeros((1, 1024), dtype=np.float32)
79
+ return sam_level0, sam_level1, wd14_embedding
80
+
81
+
82
+ def colorize_onnx(
83
+ session: ort.InferenceSession,
84
+ L_bw: np.ndarray,
85
+ sam_level0: np.ndarray,
86
+ sam_level1: np.ndarray,
87
+ wd14_embedding: np.ndarray,
88
+ ) -> np.ndarray:
89
+ """Run generator ONNX inference. Returns RGB (H, W, 3) in [0, 255]."""
90
+ L_norm = (L_bw.astype(np.float32) / 127.5) - 1.0
91
+ L_tensor = L_norm[np.newaxis, np.newaxis, :, :] # (1, 1, H, W)
92
+
93
+ ort_inputs = {
94
+ "L_bw": L_tensor,
95
+ "sam_level0": sam_level0,
96
+ "sam_level1": sam_level1,
97
+ "wd14_embedding": wd14_embedding,
98
+ }
99
+
100
+ rgb_pred = session.run(None, ort_inputs)[0] # (1, 3, H, W)
101
+ rgb_pred = rgb_pred[0].transpose(1, 2, 0) # (H, W, 3)
102
+ return denormalize_rgb(rgb_pred)
103
+
104
+
105
+ # ============================================================================
106
+ # MODEL LOADING (once, at startup)
107
+ # ============================================================================
108
+
109
+ def load_sessions():
110
+ """Load generator (+ optional SAM) ONNX sessions. Prefers CUDA if available."""
111
+ available = ort.get_available_providers()
112
+ if "CUDAExecutionProvider" in available:
113
+ providers = ["CUDAExecutionProvider", "CPUExecutionProvider"]
114
+ else:
115
+ providers = ["CPUExecutionProvider"]
116
+
117
+ if not GENERATOR_PATH.exists():
118
+ raise FileNotFoundError(f"Generator ONNX not found: {GENERATOR_PATH}")
119
+
120
+ print(f"[startup] Loading generator: {GENERATOR_PATH}", flush=True)
121
+ session = ort.InferenceSession(str(GENERATOR_PATH), providers=providers)
122
+ print(f"[startup] Generator provider: {session.get_providers()[0]}", flush=True)
123
+
124
+ sam_session = None
125
+ if SAM_PATH.exists():
126
+ print(f"[startup] Loading SAM encoder: {SAM_PATH}", flush=True)
127
+ sam_session = ort.InferenceSession(str(SAM_PATH), providers=providers)
128
+ print("[startup] SAM encoder loaded", flush=True)
129
+ else:
130
+ print("[startup] SAM encoder NOT found -> using zeros", flush=True)
131
+
132
+ return session, sam_session
133
+
134
+
135
+ SESSION, SAM_SESSION = load_sessions()
136
+ HAS_SAM = SAM_SESSION is not None
137
+
138
+
139
+ # ============================================================================
140
+ # GRADIO INFERENCE HANDLER
141
+ # ============================================================================
142
+
143
+ def colorize_image(input_image: Image.Image, infer_size: int):
144
+ """
145
+ Colorize a grayscale manga image using local ONNX models.
146
+
147
+ Args:
148
+ input_image: PIL Image (any mode).
149
+ infer_size: Square inference resolution.
150
+
151
+ Returns:
152
+ (colorized PIL Image or None, status message).
153
+ """
154
+ if input_image is None:
155
+ return None, "⚠️ Please upload an image first."
156
+
157
+ t_start = time.time()
158
+
159
+ # PIL -> grayscale numpy
160
+ gray = np.array(input_image.convert("L"))
161
+ orig_H, orig_W = gray.shape
162
+
163
+ infer_size = int(infer_size)
164
+
165
+ # Always resize input to infer_size for inference
166
+ L_bw = cv2.resize(gray, (infer_size, infer_size), interpolation=cv2.INTER_AREA)
167
+ H_in, W_in = L_bw.shape
168
+ L_norm = (L_bw.astype(np.float32) / 127.5) - 1.0
169
+
170
+ if HAS_SAM:
171
+ sam_level0, sam_level1, wd14_embedding = extract_sam_features_onnx(SAM_SESSION, L_norm)
172
+ else:
173
+ sam_level0 = np.zeros((1, 256, H_in // 16, W_in // 16), dtype=np.float32)
174
+ sam_level1 = np.zeros((1, 256, H_in // 32, W_in // 32), dtype=np.float32)
175
+ wd14_embedding = np.zeros((1, 1024), dtype=np.float32)
176
+
177
+ rgb_output = colorize_onnx(SESSION, L_bw, sam_level0, sam_level1, wd14_embedding)
178
+ # (infer_size, infer_size, 3) -> back to original input resolution
179
+ rgb_output = cv2.resize(rgb_output, (orig_W, orig_H), interpolation=cv2.INTER_LANCZOS4)
180
+
181
+ result = Image.fromarray(rgb_output)
182
+ elapsed = time.time() - t_start
183
+ status = (
184
+ f"βœ… Colorization complete! "
185
+ f"({orig_W}Γ—{orig_H} px, infer {infer_size}Γ—{infer_size}, {elapsed:.2f}s)"
186
+ )
187
+ return result, status
188
+
189
+
190
+ # ============================================================================
191
+ # GRADIO UI
192
+ # ============================================================================
193
+
194
+ def collect_examples():
195
+ """Build example list from the input/ folder."""
196
+ examples = []
197
+ if EXAMPLES_DIR.is_dir():
198
+ for ext in ("*.jpg", "*.jpeg", "*.png", "*.bmp", "*.webp"):
199
+ for f in sorted(EXAMPLES_DIR.glob(ext)):
200
+ examples.append([str(f), DEFAULT_INFER_SIZE])
201
+ return examples
202
+
203
+
204
+ def build_interface() -> gr.Blocks:
205
+ with gr.Blocks(
206
+ title="Manga Light Colorizer",
207
+ theme=gr.themes.Soft(),
208
+ ) as demo:
209
+ gr.Markdown(
210
+ """
211
+ # 🎨 Manga Light Colorizer
212
+ Upload a black-and-white manga image and let the AI bring it to life in color.
213
+
214
+ > Runs **fully locally** with ONNX Runtime β€” no external API call.
215
+ > The model was trained at **512Γ—512**; the further the inference resolution
216
+ > differs from 512, the less faithful the colors may be.
217
+ """
218
+ )
219
+
220
+ with gr.Row():
221
+ with gr.Column(scale=1):
222
+ input_image = gr.Image(label="Input Image", type="pil")
223
+ infer_size = gr.Radio(
224
+ choices=INFER_SIZE_OPTIONS,
225
+ value=DEFAULT_INFER_SIZE,
226
+ label="Inference Resolution",
227
+ info=(
228
+ "Square resolution used for inference. Output is resized back "
229
+ "to the original input resolution. 512 = best color fidelity."
230
+ ),
231
+ )
232
+ colorize_btn = gr.Button("🎨 Colorize", variant="primary", size="lg")
233
+
234
+ with gr.Column(scale=1):
235
+ output_image = gr.Image(
236
+ label="Colorized Output",
237
+ type="pil",
238
+ interactive=False,
239
+ )
240
+ status_text = gr.Textbox(label="Status", interactive=False, lines=2)
241
+
242
+ colorize_btn.click(
243
+ fn=colorize_image,
244
+ inputs=[input_image, infer_size],
245
+ outputs=[output_image, status_text],
246
+ )
247
+
248
+ examples = collect_examples()
249
+ if examples:
250
+ gr.Examples(
251
+ examples=examples,
252
+ inputs=[input_image, infer_size],
253
+ outputs=[output_image, status_text],
254
+ fn=colorize_image,
255
+ cache_examples=False,
256
+ )
257
+
258
+ gr.Markdown(
259
+ """
260
+ ---
261
+ ### πŸ“ Notes
262
+ - Supported input formats: **JPEG, PNG, WebP, BMP**.
263
+ - Inference runs locally via **ONNX Runtime** (CUDA if available, else CPU).
264
+ - Pipeline: `grayscale β†’ resize β†’ SAM encoder β†’ generator β†’ resize to original`.
265
+ """
266
+ )
267
+
268
+ return demo
269
+
270
+
271
+ print("[startup] calling build_interface()...", flush=True)
272
+ demo = build_interface()
273
+ print(f"[startup] demo object created: {demo}", flush=True)
274
+
275
+ if __name__ == "__main__":
276
+ print("[startup] running as __main__, calling demo.launch()", flush=True)
277
+ demo.launch()
packages.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ python3-opencv
requirements-cuda.txt CHANGED
@@ -3,3 +3,7 @@
3
  onnxruntime-gpu>=1.16.0
4
  numpy>=1.24.0
5
  opencv-python>=4.8.0
 
 
 
 
 
3
  onnxruntime-gpu>=1.16.0
4
  numpy>=1.24.0
5
  opencv-python>=4.8.0
6
+
7
+ # Gradio app (app.py) β€” local ONNX inference on Hugging Face Spaces
8
+ gradio>=6.0.0
9
+ Pillow>=10.0.0
requirements.txt CHANGED
@@ -3,3 +3,8 @@
3
  onnxruntime>=1.16.0
4
  numpy>=1.24.0
5
  opencv-python>=4.8.0
 
 
 
 
 
 
3
  onnxruntime>=1.16.0
4
  numpy>=1.24.0
5
  opencv-python>=4.8.0
6
+
7
+ # Gradio app (app.py) β€” local ONNX inference on Hugging Face Spaces
8
+ # On HF Spaces, OpenCV also needs the Debian package python3-opencv (see packages.txt).
9
+ gradio>=6.0.0
10
+ Pillow>=10.0.0