ImageDenosing / fastDVDnet /python /onnx_video_infer.py
wzf19947's picture
first commit
1266aec
Raw
History Blame Contribute Delete
6.64 kB
#!/usr/bin/env python
"""
FastDVDnet ONNX 视频推理脚本,输出 原帧 | ONNX 去噪结果 拼接视频。
输入视频逐帧 resize 到 ONNX 固定输入尺寸 (HxW),组成 5 帧窗口
[t-2,t-1,t,t+1,t+2] 推理,输出 resize 回原始尺寸生成拼接视频。
用法:
python onnx_video_infer.py \
--onnx ./fastdvdnet_640x480.onnx \
--video mp4/drone.mp4 \
--noise_sigma 25 \
--out_dir ./video_infer_results/onnx_test
"""
import argparse
import os
import time
import cv2
import numpy as np
import onnxruntime as ort
def get_onnx_shapes(onnx_path):
import onnx
m = onnx.load(onnx_path)
x_dims = [d.dim_value for d in m.graph.input[0].type.tensor_type.shape.dim]
return x_dims # [N, 15, H, W]
def read_video(video_path, max_frames=0):
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
raise RuntimeError("failed to open video: {}".format(video_path))
fps = cap.get(cv2.CAP_PROP_FPS) or 25.0
orig_w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
orig_h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
frames_bgr = []
while True:
ok, frame = cap.read()
if not ok:
break
frames_bgr.append(frame)
if max_frames and len(frames_bgr) >= max_frames:
break
cap.release()
if not frames_bgr:
raise RuntimeError("no frames read from {}".format(video_path))
return frames_bgr, fps, orig_w, orig_h
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_onnx_input(frame_bgr, onnx_h, onnx_w):
"""resize BGR uint8 to ONNX RGB float32 [1,3,H,W] in [0,1]"""
resized = cv2.resize(frame_bgr, (onnx_w, onnx_h), interpolation=cv2.INTER_AREA)
rgb = cv2.cvtColor(resized, cv2.COLOR_BGR2RGB)
chw = rgb.astype(np.float32).transpose(2, 0, 1) / 255.0
return chw # [3, H, W]
def chw_to_bgr_uint8(chw, target_w, target_h):
"""[3,H,W] float32 in [0,1] -> BGR uint8 resized to target_w x target_h"""
hwc = (chw * 255.0).clip(0, 255).astype(np.uint8).transpose(1, 2, 0)
rgb = hwc # already RGB from ONNX output
bgr = cv2.cvtColor(rgb, cv2.COLOR_RGB2BGR)
if bgr.shape[1] != target_w or bgr.shape[0] != target_h:
bgr = cv2.resize(bgr, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
return bgr
def denoise_onnx(frames_bgr, noise_sigma_01, sess, x_shape, input_names, output_names):
numframes = len(frames_bgr)
_, _, onnx_h, onnx_w = x_shape
temp_psz, ctrl = 5, 2
# 缓存 resize 后的 CHW 帧
chw_cache = {}
def get_chw(i):
i = i % numframes
if i not in chw_cache:
chw_cache[i] = bgr_to_onnx_input(frames_bgr[reflect_index(i, numframes)], onnx_h, onnx_w)
return chw_cache[i]
den_frames_bgr = []
inframes = [] # numpy [3,H,W] each
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))
# concat 5 frames -> [1, 15, H, W]
noisy = np.concatenate(inframes, axis=0)[None, :, :, :].astype(np.float32)
noise_map = np.full((1, 1, onnx_h, onnx_w), noise_sigma_01, dtype=np.float32)
feeds = {input_names[0]: noisy, input_names[1]: noise_map}
out = sess.run(output_names, feeds)[0] # [1, 3, H, W]
out = np.clip(out, 0.0, 1.0)
den_bgr = chw_to_bgr_uint8(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)
fourcc = cv2.VideoWriter_fourcc(*"mp4v")
writer = cv2.VideoWriter(out_path, fourcc, fps, (w * 2, h))
if not writer.isOpened():
raise RuntimeError("failed to create video writer: {}".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, "ONNX Denoised", (w + 16, 34), cv2.FONT_HERSHEY_SIMPLEX, 1.0, (0, 255, 255), 2, cv2.LINE_AA)
writer.write(canvas)
writer.release()
def main():
parser = argparse.ArgumentParser(description="FastDVDnet ONNX video inference side-by-side")
parser.add_argument("--onnx", type=str, default="./fastdvdnet_640x480.onnx")
parser.add_argument("--video", type=str, required=True, help="input mp4 video")
parser.add_argument("--noise_sigma", type=float, default=25.0, help="noise sigma 0-255")
parser.add_argument("--out_dir", type=str, default="./video_infer_results/onnx_test")
parser.add_argument("--max_frames", type=int, default=0, help="0 means all frames")
parser.add_argument("--no_label", action="store_true")
args = parser.parse_args()
x_shape = get_onnx_shapes(args.onnx)
_, _, onnx_h, onnx_w = x_shape
print("ONNX input shape: {}".format(x_shape))
sess = ort.InferenceSession(args.onnx, providers=["CPUExecutionProvider"])
input_names = [i.name for i in sess.get_inputs()]
output_names = [o.name for o in sess.get_outputs()]
print("providers:", sess.get_providers())
print("inputs:", [(n, list(i.shape)) for n, i in zip(input_names, sess.get_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, "{}_onnx_side_by_side.mp4".format(base))
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("onnx input: {}x{}".format(onnx_w, onnx_h))
sigma01 = args.noise_sigma / 255.0
den = denoise_onnx(frames_bgr, sigma01, sess, x_shape, input_names, output_names)
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()