nanovsr-demo / inference.py
filippawlicki's picture
Add Gradio demo app for NanoVSR
6870c03
Raw
History Blame Contribute Delete
3.99 kB
"""Checkpoint management and inference helpers for the NanoVSR Gradio demo.
Vendored/adapted from the main nanovsr repo's utils.py + demo.py so this Space
has no dependency on the training repo at runtime.
"""
import os
import numpy as np
import requests
import torch
from models.nanovsr import NanoVSR
SCALE = 4
CKPT_DIR = os.path.join(os.path.dirname(__file__), "checkpoints")
RELEASE_BASE = "https://github.com/filippawlicki/nanovsr/releases/download/v1.0"
MODEL_INFO = {
"NanoVSR-226k (fastest)": dict(file="nanovsr_226k.pth", num_feat=32, num_blocks=8,
psnr="28.23 dB REDS4"),
"NanoVSR-644k (baseline)": dict(file="nanovsr_644k.pth", num_feat=48, num_blocks=12,
psnr="28.64 dB REDS4"),
"NanoVSR-1.7M": dict(file="nanovsr_1.7m.pth", num_feat=64, num_blocks=20,
psnr="29.15 dB REDS4"),
"NanoVSR-5.4M (best quality)": dict(file="nanovsr_5.4m.pth", num_feat=96, num_blocks=30,
psnr="29.73 dB REDS4"),
}
DEFAULT_MODEL = "NanoVSR-644k (baseline)"
_model_cache = {}
def ensure_checkpoint(model_name):
info = MODEL_INFO[model_name]
os.makedirs(CKPT_DIR, exist_ok=True)
path = os.path.join(CKPT_DIR, info["file"])
if not os.path.exists(path):
url = f"{RELEASE_BASE}/{info['file']}"
tmp_path = path + ".part"
with requests.get(url, stream=True, timeout=60) as resp:
resp.raise_for_status()
with open(tmp_path, "wb") as f:
for chunk in resp.iter_content(chunk_size=1 << 20):
f.write(chunk)
os.replace(tmp_path, path)
return path
def get_model(model_name):
"""Load (and cache on CPU) the requested NanoVSR variant."""
if model_name in _model_cache:
return _model_cache[model_name]
info = MODEL_INFO[model_name]
checkpoint_path = ensure_checkpoint(model_name)
state_dict = torch.load(checkpoint_path, map_location="cpu")
if isinstance(state_dict, dict):
for key in ["params_ema", "params", "model_state_dict"]:
if key in state_dict:
state_dict = state_dict[key]
break
model = NanoVSR(num_feat=info["num_feat"], num_blocks=info["num_blocks"])
model.load_state_dict(state_dict, strict=False)
model.switch_to_deploy()
model.eval()
_model_cache[model_name] = model
return model
def frames_to_tensor(frames_rgb):
"""List/array of [H, W, 3] uint8 RGB frames -> [T, 3, H, W] float tensor in [0, 1]."""
batch = np.stack(frames_rgb, axis=0).transpose(0, 3, 1, 2)
batch = np.ascontiguousarray(batch)
return torch.from_numpy(batch).float() / 255.0
def tensor_to_frames(tensor):
"""[T, 3, H, W] float tensor in [0, 1] -> [T, H, W, 3] uint8 RGB numpy array."""
arr = (tensor.clamp(0, 1).mul(255.0).round().to(torch.uint8)
.permute(0, 2, 3, 1).contiguous().cpu().numpy())
return arr
@torch.inference_mode()
def run_chunked_inference(model, frames_rgb, chunk_size, device, use_fp16, progress_cb=None):
"""Run NanoVSR over a list of LR RGB frames in temporal chunks.
Returns a [T, H*4, W*4, 3] uint8 RGB numpy array.
"""
model = model.to(device)
autocast_ctx = (torch.autocast(device_type="cuda", dtype=torch.float16)
if use_fp16 and device == "cuda" else torch.autocast(device_type="cpu", enabled=False))
outputs = []
n = len(frames_rgb)
for start in range(0, n, chunk_size):
chunk = frames_rgb[start:start + chunk_size]
lq = frames_to_tensor(chunk).unsqueeze(0).to(device)
with autocast_ctx:
sr = model(lq)
outputs.append(tensor_to_frames(sr.float().squeeze(0)))
if progress_cb is not None:
progress_cb(min(start + chunk_size, n), n)
model.cpu()
if device == "cuda":
torch.cuda.empty_cache()
return np.concatenate(outputs, axis=0)