复制
import cv2
import torch
from model.RIFE import Model
model = Model()
model.load_model('./train_log', -1)
model.eval()
model.device()
def process_video(input_path, output_path, multiplier=2):
cap = cv2.VideoCapture(input_path)
fps = cap.get(cv2.CAP_PROP_FPS)
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
out = cv2.VideoWriter(
output_path,
cv2.VideoWriter_fourcc(*'mp4v'),
fps * multiplier,
(width, height)
)
ret, prev_frame = cap.read()
if not ret:
return
out.write(prev_frame)
while True:
ret, curr_frame = cap.read()
if not ret:
break
# 准备张量
img0 = torch.from_numpy(prev_frame).permute(2, 0, 1).float() / 255.0
img1 = torch.from_numpy(curr_frame).permute(2, 0, 1).float() / 255.0
img0 = img0.unsqueeze(0).cuda()
img1 = img1.unsqueeze(0).cuda()
# 生成中间帧
for i in range(multiplier - 1):
t = (i + 1) / multiplier
with torch.no_grad():
middle = model.inference(img0, img1, timestep=t)
middle = (middle[0] * 255).byte().cpu().numpy().transpose(1, 2, 0)
out.write(middle)
out.write(curr_frame)
prev_frame = curr_frame
cap.release()
out.release()
process_video('input.mp4', 'output_60fps.mp4', multiplier=2)