ImageDenosing / fastDVDnet /python /axmodel_video_infer.py
wzf19947's picture
update
9d449ad
Raw
History Blame Contribute Delete
8.11 kB
#!/usr/bin/env python
"""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():
# fallback to mp4v if avc1 not available
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()