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()