| |
| """FastDVDnet axmodel video inference, side-by-side output. |
| Input: mp4 video -> resize to axmodel fixed input -> 5-frame window -> axmodel -> output video. |
| Input dtype is inferred from the axmodel session so both float and uint8 models work.""" |
| import argparse, os, time |
| import cv2, numpy as np |
| import axengine as axe |
|
|
|
|
| def _resolve_numpy_dtype(dtype, fallback): |
| try: |
| return np.dtype(dtype) |
| except TypeError: |
| name = str(dtype).lower() |
| if "float16" in name or "fp16" in name: |
| return np.dtype(np.float16) |
| if "float" in name: |
| return np.dtype(np.float32) |
| if "uint8" in name: |
| return np.dtype(np.uint8) |
| if "int8" in name: |
| return np.dtype(np.int8) |
| return np.dtype(fallback) |
|
|
| def read_video(video_path, max_frames=0): |
| cap=cv2.VideoCapture(video_path) |
| if not cap.isOpened(): raise RuntimeError('open fail: '+video_path) |
| fps=cap.get(cv2.CAP_PROP_FPS) or 25.0 |
| ow=int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)); oh=int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) |
| frames=[] |
| while True: |
| ok,f=cap.read() |
| if not ok: break |
| frames.append(f) |
| if max_frames and len(frames)>=max_frames: break |
| cap.release() |
| if not frames: raise RuntimeError('no frames') |
| return frames,fps,ow,oh |
|
|
| def reflect_index(idx,length): |
| if length<=1: return 0 |
| while idx<0 or idx>=length: |
| if idx<0: idx=-idx |
| if idx>=length: idx=2*(length-1)-idx |
| return idx |
|
|
| def bgr_to_ax_input(fb, ah, aw, noisy_dtype): |
| r = cv2.resize(fb, (aw, ah), interpolation=cv2.INTER_AREA) |
| r = cv2.cvtColor(r, cv2.COLOR_BGR2RGB) |
| chw = r.transpose(2, 0, 1) |
| noisy_dtype = _resolve_numpy_dtype(noisy_dtype, np.float32) |
| if np.issubdtype(noisy_dtype, np.floating): |
| return (chw.astype(np.float32) / 255.0).astype(noisy_dtype, copy=False) |
| return chw.astype(noisy_dtype, copy=False) |
|
|
| def chw_float_to_bgr_u8(chw,tw,th): |
| hwc=(chw*255.).clip(0,255).astype(np.uint8).transpose(1,2,0) |
| bgr=cv2.cvtColor(hwc,cv2.COLOR_RGB2BGR) |
| if bgr.shape[1]!=tw or bgr.shape[0]!=th: |
| bgr=cv2.resize(bgr,(tw,th),interpolation=cv2.INTER_LINEAR) |
| return bgr |
|
|
| def denoise_axmodel(frames_bgr, noise_sigma_01, sess, ax_h, ax_w, |
| input_names, output_names, noisy_dtype, noise_map_dtype): |
| numframes = len(frames_bgr) |
| temp_psz, ctrl = 5, 2 |
|
|
| chw_cache = {} |
| def get_chw(i): |
| i = i % numframes |
| if i not in chw_cache: |
| chw_cache[i] = bgr_to_ax_input( |
| frames_bgr[reflect_index(i, numframes)], ax_h, ax_w, noisy_dtype) |
| return chw_cache[i] |
|
|
| den_frames_bgr = [] |
| inframes = [] |
|
|
| for fridx in range(numframes): |
| if not inframes: |
| for off in range(temp_psz): |
| inframes.append(get_chw(fridx + off - ctrl)) |
| else: |
| del inframes[0] |
| inframes.append(get_chw(fridx + ctrl)) |
|
|
| noisy = np.concatenate(inframes, axis=0)[None, :, :, :] |
| noise_map = np.full( |
| (1, 1, ax_h, ax_w), |
| noise_sigma_01, |
| dtype=_resolve_numpy_dtype(noise_map_dtype, np.float32), |
| ) |
| feeds = {input_names[0]: noisy, input_names[1]: noise_map} |
| out = sess.run(output_names, feeds)[0] |
| out = np.clip(out, 0.0, 1.0) |
| den_bgr = chw_float_to_bgr_u8( |
| out[0], frames_bgr[0].shape[1], frames_bgr[0].shape[0]) |
| den_frames_bgr.append(den_bgr) |
|
|
| return den_frames_bgr |
|
|
|
|
| def write_side_by_side(frames_orig, frames_den, out_path, fps, label=True): |
| h, w = frames_orig[0].shape[:2] |
| os.makedirs(os.path.dirname(os.path.abspath(out_path)), exist_ok=True) |
|
|
| ext = os.path.splitext(out_path)[1].lower() |
|
|
| if ext == '.gif': |
| _write_gif(frames_orig, frames_den, out_path, fps, w, h, label) |
| else: |
| _write_mp4(frames_orig, frames_den, out_path, fps, w, h, label) |
|
|
|
|
| def _write_mp4(frames_orig, frames_den, out_path, fps, w, h, label): |
| fourcc = cv2.VideoWriter_fourcc(*"avc1") |
| writer = cv2.VideoWriter(out_path, fourcc, fps, (w * 2, h)) |
| if not writer.isOpened(): |
| |
| fourcc = cv2.VideoWriter_fourcc(*"mp4v") |
| writer = cv2.VideoWriter(out_path, fourcc, fps, (w * 2, h)) |
| if not writer.isOpened(): |
| raise RuntimeError("video writer fail: {}".format(out_path)) |
| for orig, den in zip(frames_orig, frames_den): |
| canvas = np.concatenate([orig, den], axis=1) |
| if label: |
| cv2.putText(canvas, "Original", (16, 34), cv2.FONT_HERSHEY_SIMPLEX, |
| 1.0, (0, 255, 255), 2, cv2.LINE_AA) |
| cv2.putText(canvas, "AXModel Denoised", (w + 16, 34), cv2.FONT_HERSHEY_SIMPLEX, |
| 1.0, (0, 255, 255), 2, cv2.LINE_AA) |
| writer.write(canvas) |
| writer.release() |
|
|
|
|
| def _write_gif(frames_orig, frames_den, out_path, fps, w, h, label): |
| try: |
| from PIL import Image |
| except ImportError: |
| raise RuntimeError("GIF output requires Pillow: pip install Pillow") |
|
|
| gif_frames = [] |
| for orig, den in zip(frames_orig, frames_den): |
| canvas = np.concatenate([orig, den], axis=1) |
| if label: |
| cv2.putText(canvas, "Original", (16, 34), cv2.FONT_HERSHEY_SIMPLEX, |
| 1.0, (0, 255, 255), 2, cv2.LINE_AA) |
| cv2.putText(canvas, "AXModel Denoised", (w + 16, 34), cv2.FONT_HERSHEY_SIMPLEX, |
| 1.0, (0, 255, 255), 2, cv2.LINE_AA) |
| canvas_rgb = cv2.cvtColor(canvas, cv2.COLOR_BGR2RGB) |
| gif_frames.append(Image.fromarray(canvas_rgb)) |
|
|
| duration = int(1000.0 / fps) |
| gif_frames[0].save( |
| out_path, save_all=True, append_images=gif_frames[1:], |
| duration=duration, loop=0, optimize=True, |
| ) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="FastDVDnet axmodel video inference") |
| parser.add_argument("--axmodel", type=str, default="./fastdvdnet_640x480.axmodel") |
| parser.add_argument("--video", type=str, default='./people-sunset.mp4', help="input mp4 video") |
| parser.add_argument("--noise_sigma", type=float, default=25.0, help="0-255") |
| parser.add_argument("--out_dir", type=str, default="./") |
| parser.add_argument("--max_frames", type=int, default=0, help="0=all") |
| parser.add_argument("--no_label", action="store_true") |
| parser.add_argument("--format", type=str, default="gif", choices=["mp4", "gif"], |
| help="output format") |
| args = parser.parse_args() |
|
|
| sess = axe.InferenceSession(args.axmodel, providers=["AxEngineExecutionProvider"]) |
| inputs = sess.get_inputs() |
| input_names = [i.name for i in inputs] |
| output_names = [o.name for o in sess.get_outputs()] |
| noisy_dtype = getattr(inputs[0], "dtype", np.float32) |
| noise_map_dtype = getattr(inputs[1], "dtype", np.float32) |
| ax_h = inputs[0].shape[2] |
| ax_w = inputs[0].shape[3] |
|
|
| print("providers:", sess.get_providers()) |
| print("inputs:", [(n, list(i.shape), str(getattr(i, "dtype", "unknown"))) for n, i in zip(input_names, inputs)]) |
| print("outputs:", [(n, list(o.shape)) for n, o in zip(output_names, sess.get_outputs())]) |
|
|
| base = os.path.splitext(os.path.basename(args.video))[0] |
| out_path = os.path.join(args.out_dir, base + "_axmodel_side_by_side." + args.format) |
|
|
| t0 = time.time() |
| frames_bgr, fps, orig_w, orig_h = read_video(args.video, args.max_frames) |
| print("video: {}x{} @ {:.1f}fps, {} frames".format(orig_w, orig_h, fps, len(frames_bgr))) |
| print("axmodel input: {}x{}".format(ax_w, ax_h)) |
|
|
| sigma01 = args.noise_sigma / 255.0 |
| den = denoise_axmodel( |
| frames_bgr, sigma01, sess, ax_h, ax_w, |
| input_names, output_names, noisy_dtype, noise_map_dtype, |
| ) |
| write_side_by_side(frames_bgr, den, out_path, fps, label=not args.no_label) |
| dt = time.time() - t0 |
|
|
| print("[OK] {} -> {}".format(args.video, out_path)) |
| print(" frames={} time={:.2f}s fps={:.2f}".format( |
| len(frames_bgr), dt, len(frames_bgr) / max(dt, 0.001))) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|