File size: 6,635 Bytes
1266aec | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 | #!/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()
|