Spaces:
Running
Running
| <html lang="en"> | |
| <head> | |
| <meta charset="utf-8" /> | |
| <meta name="viewport" content="width=device-width, initial-scale=1" /> | |
| <title>AI Tiled Upscaler — WebGPU</title> | |
| <!-- FIX 1: Using the official patched AWS endpoints to avoid the 4.44.1 CDN crash --> | |
| <script type="module" crossorigin | |
| src="https://gradio-lite-previews.s3.amazonaws.com/PINNED_HF_HUB/dist/lite.js"> | |
| </script> | |
| <link rel="stylesheet" | |
| href="https://gradio-lite-previews.s3.amazonaws.com/PINNED_HF_HUB/dist/lite.css" /> | |
| <style> | |
| *, *::before, *::after { box-sizing: border-box; margin: 0; padding: 0; } | |
| :root { | |
| --bg: #0a0a0f; | |
| --bg2: #12121a; | |
| --accent: #7c6aff; | |
| --accent2: #b8aaff; | |
| --text: #e8e6ff; | |
| --muted: #7a7898; | |
| --border: rgba(124,106,255,0.18); | |
| } | |
| html, body { | |
| height: 100%; | |
| background: var(--bg); | |
| color: var(--text); | |
| font-family: 'DM Sans', system-ui, sans-serif; | |
| overflow-x: hidden; | |
| } | |
| body::before { | |
| content: ''; | |
| position: fixed; inset: 0; z-index: 0; | |
| background-image: | |
| linear-gradient(rgba(124,106,255,0.04) 1px, transparent 1px), | |
| linear-gradient(90deg, rgba(124,106,255,0.04) 1px, transparent 1px); | |
| background-size: 40px 40px; | |
| pointer-events: none; | |
| } | |
| header { | |
| position: relative; z-index: 1; | |
| padding: 28px 32px 0; | |
| display: flex; align-items: baseline; gap: 14px; | |
| } | |
| header h1 { | |
| font-size: 22px; font-weight: 700; | |
| letter-spacing: -0.03em; color: var(--accent2); | |
| } | |
| .badge { | |
| font-family: 'JetBrains Mono', monospace; | |
| font-size: 10px; font-weight: 500; | |
| padding: 2px 8px; border-radius: 4px; | |
| background: rgba(124,106,255,0.15); | |
| border: 1px solid var(--border); | |
| color: var(--accent2); | |
| letter-spacing: 0.06em; text-transform: uppercase; | |
| } | |
| #model-status { | |
| position: relative; z-index: 1; | |
| margin: 10px 32px 0; | |
| font-family: 'JetBrains Mono', monospace; | |
| font-size: 11px; color: var(--muted); height: 16px; | |
| transition: color 0.3s; | |
| } | |
| #model-status.ready { color: #5dd78a; } | |
| #model-status.error { color: #f06060; } | |
| #gradio-wrap { | |
| position: relative; z-index: 1; | |
| margin: 18px 24px 24px; | |
| border-radius: 12px; overflow: hidden; | |
| border: 1px solid var(--border); | |
| background: var(--bg2); | |
| min-height: 400px; | |
| } | |
| gradio-lite { | |
| --color-accent: #7c6aff ; | |
| --body-background-fill: #12121a ; | |
| --block-background-fill: #1a1a26 ; | |
| --border-color-primary: rgba(124,106,255,0.2) ; | |
| --button-primary-background-fill: #7c6aff ; | |
| --button-primary-text-color: #fff ; | |
| } | |
| footer { | |
| position: relative; z-index: 1; | |
| padding: 0 32px 28px; | |
| font-size: 11px; color: var(--muted); line-height: 1.7; | |
| } | |
| footer a { color: var(--accent2); text-decoration: none; } | |
| footer a:hover { text-decoration: underline; } | |
| @import url('https://fonts.googleapis.com/css2?family=DM+Sans:wght@400;500;700&family=JetBrains+Mono:wght@400;500&display=swap'); | |
| </style> | |
| </head> | |
| <body> | |
| <header> | |
| <h1>AI Tiled Upscaler</h1> | |
| <span class="badge">WebGPU · runs locally</span> | |
| </header> | |
| <div id="model-status">Initialising model…</div> | |
| <div id="gradio-wrap"> | |
| <!-- Shell div: Gradio-Lite will be injected here dynamically --> | |
| </div> | |
| <footer> | |
| Runs entirely in your browser — no image is uploaded anywhere. | | |
| Model: <a href="https://huggingface.co/Xenova/swin2SR-compressed-sr-x4-48" target="_blank">swin2SR-compressed-sr-x4-48</a> | |
| via <a href="https://github.com/xenova/transformers.js" target="_blank">Transformers.js v2</a>. | |
| </footer> | |
| <script type="module"> | |
| const APP_PY = ` | |
| import gradio as gr | |
| from core import ( | |
| slice_image_to_memory, | |
| upscale_all_tiles, | |
| stitch_from_memory, | |
| tiles_to_zip_b64, | |
| UPSCALE_FACTOR, | |
| MAX_INPUT_PIXELS, | |
| pil_to_b64, | |
| ) | |
| MAX_SIDE = int(MAX_INPUT_PIXELS ** 0.5) | |
| async def handle_upload(input_image, progress=gr.Progress(track_tqdm=False)): | |
| if input_image is None: | |
| raise gr.Error("Please upload an image first.") | |
| try: | |
| import asyncio | |
| progress(0, desc="Phase 1: Slicing image...") | |
| await asyncio.sleep(0.1) | |
| canvas_w, canvas_h, tile_w, tile_h, tiles, pad_left, pad_top, orig_w, orig_h = slice_image_to_memory(input_image) | |
| await upscale_all_tiles(tiles, progress=progress) | |
| progress(0.90, desc="Phase 2: Upscaling finished — preparing tile ZIP download...") | |
| await asyncio.sleep(0.1) | |
| zip_b64 = tiles_to_zip_b64(tiles) | |
| download_html = ( | |
| '<a href="data:application/zip;base64,' + zip_b64 + '" ' | |
| 'download="upscaled_tiles.zip" ' | |
| 'style="display:inline-block;padding:12px 20px;background:#7c6aff;color:#ffffff;' | |
| 'border-radius:10px;text-decoration:none;font-weight:700;">' | |
| '📥 Download Upscaled Tiles (ZIP)</a>' | |
| ) | |
| yield [download_html, None] | |
| await asyncio.sleep(0.1) | |
| progress(0.92, desc="Phase 3: Stitching tiles (heavy math, please wait)...") | |
| await asyncio.sleep(0.1) | |
| final_image = stitch_from_memory( | |
| canvas_w * UPSCALE_FACTOR, | |
| canvas_h * UPSCALE_FACTOR, | |
| tile_w * UPSCALE_FACTOR, | |
| tile_h * UPSCALE_FACTOR, | |
| tiles, | |
| UPSCALE_FACTOR, | |
| pad_left, | |
| pad_top, | |
| orig_w, | |
| orig_h, | |
| ) | |
| progress(0.98, desc="Phase 4: Encoding final image...") | |
| await asyncio.sleep(0.1) | |
| final_b64 = "data:image/jpeg;base64," + await pil_to_b64(final_image, format="JPEG") | |
| yield [download_html, final_b64] | |
| except ValueError as ve: | |
| raise gr.Error(str(ve)) | |
| except Exception as ex: | |
| raise gr.Error("Upscaling failed: " + str(ex)) | |
| with gr.Blocks( | |
| title="AI Tiled Upscaler", | |
| css=".gradio-container { background: transparent !important; } footer { display: none !important; }", | |
| ) as demo: | |
| gr.Markdown( | |
| "Upload an image and hit **Upscale**. " | |
| "Each tile is processed **locally on your GPU** via WebGPU. " | |
| "No data is sent to any server. " | |
| "Output is **" + str(UPSCALE_FACTOR) + "x** the input resolution. " | |
| "Max input: **" + str(MAX_SIDE) + "x" + str(MAX_SIDE) + " px**." | |
| ) | |
| with gr.Row(): | |
| inp = gr.Image(label="Input image", type="pil", sources=["upload"]) | |
| download_link = gr.HTML(label="Tile Download") | |
| out = gr.Image(label="Upscaled output (x" + str(UPSCALE_FACTOR) + ")") | |
| btn = gr.Button("Upscale", variant="primary") | |
| btn.click(fn=handle_upload, inputs=inp, outputs=[download_link, out]) | |
| demo.launch() | |
| `; | |
| const CORE_PY = ` | |
| import io | |
| import base64 | |
| import math | |
| import zipfile | |
| import numpy as np | |
| from PIL import Image | |
| UPSCALE_FACTOR = 4 | |
| UPSCALE_MODEL = "Xenova/swin2SR-compressed-sr-x4-48" | |
| TILE_MIN = 256 | |
| TILE_MAX = 1024 | |
| TILE_SNAP = 64 | |
| TILE_TARGET_PCT = 0.55 | |
| OVERLAP_MIN_PCT = 0.10 | |
| MAX_INPUT_PIXELS = 2000 * 2000 | |
| _upscale_pipe = None | |
| async def _get_pipeline(): | |
| global _upscale_pipe | |
| if _upscale_pipe is None: | |
| from transformers_js_py import pipeline | |
| _upscale_pipe = await pipeline("image-to-image", UPSCALE_MODEL, {"dtype": "fp32", "device": "webgpu"}) | |
| return _upscale_pipe | |
| async def _canvas_encode_image(img, mime_type="image/jpeg", quality=0.95): | |
| import js | |
| if not hasattr(js, "OffscreenCanvas") or not hasattr(js, "FileReader"): | |
| return None | |
| rgba = np.asarray(img.convert("RGBA"), dtype=np.uint8) | |
| h, w = rgba.shape[:2] | |
| arr = js.Uint8ClampedArray.new(memoryview(rgba)) | |
| image_data = js.ImageData.new(arr, w, h) | |
| canvas = js.OffscreenCanvas.new(w, h) | |
| ctx = canvas.getContext("2d") | |
| ctx.putImageData(image_data, 0, 0) | |
| opts = {"type": mime_type, "quality": quality} | |
| blob = await canvas.convertToBlob(opts) | |
| reader = js.FileReader.new() | |
| promise = js.Promise.new(lambda resolve, reject: ( | |
| reader.addEventListener("load", lambda event: resolve(reader.result)), | |
| reader.addEventListener("error", lambda event: reject(event)), | |
| reader.readAsDataURL(blob) | |
| )) | |
| data_url = await promise | |
| return str(data_url).split(",", 1)[1] | |
| async def pil_to_b64(img, format="PNG"): | |
| if format == "JPEG": | |
| try: | |
| b64 = await _canvas_encode_image(img, mime_type="image/jpeg", quality=0.95) | |
| if b64 is not None: | |
| return b64 | |
| except Exception: | |
| pass | |
| buf = io.BytesIO() | |
| if format == "JPEG": | |
| img.save(buf, format=format, quality=95) | |
| else: | |
| img.save(buf, format=format) | |
| return base64.b64encode(buf.getvalue()).decode("ascii") | |
| def tiles_to_zip_b64(tiles): | |
| zip_buf = io.BytesIO() | |
| with zipfile.ZipFile(zip_buf, mode="w", compression=zipfile.ZIP_STORED) as zf: | |
| for idx, tile_entry in enumerate(tiles, 1): | |
| tile_buf = io.BytesIO() | |
| tile_entry["image"].save(tile_buf, format="PNG") | |
| zf.writestr(f"tile_{idx:03d}.png", tile_buf.getvalue()) | |
| return base64.b64encode(zip_buf.getvalue()).decode("ascii") | |
| def b64_to_pil(b64): | |
| data = base64.b64decode(b64) | |
| return Image.open(io.BytesIO(data)).convert("RGB") | |
| def choose_tile_dimension(length): | |
| target = length * TILE_TARGET_PCT | |
| snapped = (int(target) // TILE_SNAP) * TILE_SNAP | |
| return max(TILE_MIN, min(TILE_MAX, snapped)) | |
| def pad_if_needed(img, tile_w, tile_h): | |
| orig_w, orig_h = img.size | |
| pad_left = max(0, tile_w - orig_w) | |
| pad_top = max(0, tile_h - orig_h) | |
| if pad_left == 0 and pad_top == 0: | |
| return img, 0, 0, orig_w, orig_h | |
| arr = np.asarray(img.convert("RGB"), dtype=np.uint8) | |
| mode = "reflect" if pad_left < orig_w and pad_top < orig_h else "edge" | |
| arr_padded = np.pad(arr, ((pad_top, 0), (pad_left, 0), (0, 0)), mode=mode) | |
| return Image.fromarray(arr_padded, "RGB"), pad_left, pad_top, orig_w, orig_h | |
| def calc_tile_starts(image_size, tile_size): | |
| if image_size <= tile_size: | |
| return [0] | |
| max_stride = tile_size * (1.0 - OVERLAP_MIN_PCT) | |
| n = max(math.ceil((image_size - tile_size) / max_stride) + 1, 2) | |
| while True: | |
| stride = (image_size - tile_size) / (n - 1) | |
| if (tile_size - stride) / tile_size >= OVERLAP_MIN_PCT: | |
| break | |
| n += 1 | |
| starts = [] | |
| for i in range(n): | |
| c = round(i * stride) | |
| c = min(c, image_size - tile_size) | |
| starts.append(c) | |
| unique = sorted(set(starts)) | |
| assert unique[-1] + tile_size >= image_size | |
| return unique | |
| def slice_image_to_memory(pil_image): | |
| img = pil_image.convert("RGB") | |
| img_w, img_h = img.size | |
| if img_w * img_h > MAX_INPUT_PIXELS: | |
| side = int(MAX_INPUT_PIXELS ** 0.5) | |
| raise ValueError( | |
| "Image (" + str(img_w) + "x" + str(img_h) + ") exceeds the " | |
| + str(side) + "x" + str(side) + " px limit for in-browser upscaling." | |
| ) | |
| tile_w = choose_tile_dimension(img_w) | |
| tile_h = choose_tile_dimension(img_h) | |
| img, pad_left, pad_top, orig_w, orig_h = pad_if_needed(img, tile_w, tile_h) | |
| canvas_w, canvas_h = img.size | |
| xs = calc_tile_starts(canvas_w, tile_w) | |
| ys = calc_tile_starts(canvas_h, tile_h) | |
| print("Tile size: " + str(tile_w) + "x" + str(tile_h) + " Grid: " + str(len(xs)) + "x" + str(len(ys))) | |
| tiles = [] | |
| for y in ys: | |
| for x in xs: | |
| x1 = min(x, canvas_w - tile_w) | |
| y1 = min(y, canvas_h - tile_h) | |
| tiles.append({"image": img.crop((x1, y1, x1 + tile_w, y1 + tile_h)), "x": x1, "y": y1}) | |
| print(str(len(tiles)) + " tiles ready.") | |
| return canvas_w, canvas_h, tile_w, tile_h, tiles, pad_left, pad_top, orig_w, orig_h | |
| def create_cosine_bell_mask(tile_w, tile_h): | |
| ramp_x = 0.5 * (1.0 - np.cos(np.pi * np.linspace(0, 1, tile_w, dtype=np.float32))) | |
| ramp_y = 0.5 * (1.0 - np.cos(np.pi * np.linspace(0, 1, tile_h, dtype=np.float32))) | |
| return np.outer(ramp_y, ramp_x) | |
| async def upscale_tile(tile_entry): | |
| from transformers_js_py import import_transformers_js | |
| from pyodide.ffi import to_js | |
| pipe = await _get_pipeline() | |
| pil_tile = tile_entry["image"] | |
| w, h = pil_tile.size | |
| try: | |
| arr_rgb = np.asarray(pil_tile.convert("RGB"), dtype=np.uint8) | |
| js_uint8 = to_js(arr_rgb.tobytes()) | |
| transformers = await import_transformers_js() | |
| raw_image = transformers.RawImage.new(js_uint8, w, h, 3) | |
| raw_output = await pipe(raw_image) | |
| out_bytes = bytes(raw_output.data) | |
| out_arr = np.frombuffer(out_bytes, dtype=np.uint8).reshape((raw_output.height, raw_output.width, raw_output.channels)) | |
| if raw_output.channels == 4: | |
| out_arr = out_arr[:, :, :3] | |
| tile_entry["image"] = Image.fromarray(out_arr, mode="RGB") | |
| except Exception as exc: | |
| print(f"WARN: inference failed - {exc}. Using PIL bicubic fallback.") | |
| tile_entry["image"] = pil_tile.resize((w * UPSCALE_FACTOR, h * UPSCALE_FACTOR), Image.LANCZOS) | |
| async def upscale_all_tiles(tiles, progress=None): | |
| import asyncio | |
| total = len(tiles) | |
| for idx, tile_entry in enumerate(tiles, 1): | |
| desc = "Phase 2: Upscaling tile " + str(idx) + "/" + str(total) + "..." | |
| print(desc) | |
| if progress is not None: | |
| # Scale this phase between 5% and 90% of the total progress bar | |
| progress(0.05 + 0.85 * (idx / total), desc=desc) | |
| await asyncio.sleep(0.05) | |
| await upscale_tile(tile_entry) | |
| def stitch_from_memory(canvas_w, canvas_h, tile_w, tile_h, tiles, scale_factor, pad_left, pad_top, orig_w, orig_h): | |
| accumulator = np.zeros((canvas_h, canvas_w, 3), dtype=np.float32) | |
| weight_map = np.zeros((canvas_h, canvas_w), dtype=np.float32) | |
| mask = create_cosine_bell_mask(tile_w, tile_h) | |
| for tile_entry in tiles: | |
| x = tile_entry["x"] * scale_factor | |
| y = tile_entry["y"] * scale_factor | |
| tile_arr = np.asarray(tile_entry["image"].convert("RGB"), dtype=np.float32) | |
| if tile_arr.shape != (tile_h, tile_w, 3): | |
| raise RuntimeError("Tile shape mismatch at (" + str(x) + "," + str(y) + "): " + str(tile_arr.shape)) | |
| accumulator[y:y+tile_h, x:x+tile_w] += tile_arr * mask[:,:,np.newaxis] | |
| weight_map[y:y+tile_h, x:x+tile_w] += mask | |
| weight_map = weight_map[:, :, np.newaxis] | |
| final = np.clip(accumulator / (weight_map + 1e-8), 0, 255).astype(np.uint8) | |
| stitched = Image.fromarray(final, "RGB") | |
| if pad_left > 0 or pad_top > 0: | |
| cl = pad_left * scale_factor | |
| ct = pad_top * scale_factor | |
| stitched = stitched.crop((cl, ct, cl + orig_w * scale_factor, ct + orig_h * scale_factor)) | |
| return stitched | |
| async def run_upscaler_pipeline(pil_image, progress=None): | |
| canvas_w, canvas_h, tile_w, tile_h, tiles, pad_left, pad_top, orig_w, orig_h = slice_image_to_memory(pil_image) | |
| await upscale_all_tiles(tiles, progress=progress) | |
| if progress is not None: | |
| progress(0.90, desc="Phase 3: Stitching tiles (heavy math, please wait)...") | |
| import asyncio | |
| await asyncio.sleep(0.1) | |
| return stitch_from_memory( | |
| canvas_w * UPSCALE_FACTOR, canvas_h * UPSCALE_FACTOR, | |
| tile_w * UPSCALE_FACTOR, tile_h * UPSCALE_FACTOR, | |
| tiles, UPSCALE_FACTOR, | |
| pad_left, pad_top, orig_w, orig_h, | |
| ) | |
| `; | |
| async function bootGradioLite() { | |
| await customElements.whenDefined("gradio-lite"); | |
| const wrap = document.getElementById("gradio-wrap"); | |
| // Clear the existing empty wrapper content | |
| wrap.innerHTML = ""; | |
| // Dynamically create the Gradio-Lite DOM elements and inject Python via textContent. | |
| // This guarantees the browser's HTML parser will NOT touch your Python code, | |
| // preventing it from mistakenly treating '<' and '>' symbols as broken HTML tags. | |
| const gl = document.createElement("gradio-lite"); | |
| const reqs = document.createElement("gradio-requirements"); | |
| reqs.textContent = "numpy\nPillow\nhuggingface-hub==0.32.1\ntransformers_js_py"; | |
| gl.appendChild(reqs); | |
| const appFile = document.createElement("gradio-file"); | |
| appFile.setAttribute("name", "app.py"); | |
| appFile.setAttribute("entrypoint", ""); | |
| appFile.textContent = APP_PY; | |
| gl.appendChild(appFile); | |
| const coreFile = document.createElement("gradio-file"); | |
| coreFile.setAttribute("name", "core.py"); | |
| coreFile.textContent = CORE_PY; | |
| gl.appendChild(coreFile); | |
| wrap.appendChild(gl); | |
| } | |
| bootGradioLite().catch(err => { | |
| console.error("Gradio-Lite boot failed:", err); | |
| document.getElementById("model-status").textContent = "Gradio-Lite failed to load: " + err.message; | |
| document.getElementById("model-status").className = "error"; | |
| }); | |
| </script> | |
| </body> | |
| </html> |