lab3 impl

This commit is contained in:
2026-06-27 05:41:51 +03:00
commit d1e9e2f467
22 changed files with 2041 additions and 0 deletions
+30
View File
@@ -0,0 +1,30 @@
# Byte-compiled
__pycache__/
*.py[cod]
# Virtual environment
lab3_venv/
# Swarm infrastructure
.opencode/
.swarm/
# Model weights (large, re-downloadable)
rife_model/flownet.pkl
train_log/
train_log_hd/
# LaTeX build artifacts
*.aux
*.log
*.out
*.toc
*.lof
*.lot
*.bbl
*.blg
missfont.log
# Temp / generated
lab3_rife_executed.ipynb
lab3_rife_script.txt
Binary file not shown.

After

Width:  |  Height:  |  Size: 2.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 84 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 644 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 892 KiB

+678
View File
@@ -0,0 +1,678 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Лабораторная работа №3\n",
"## Нейросетевая интерполяция кадров видео с использованием RIFE\n",
"\n",
"**Дисциплина:** Разработка мультимедийных приложений \n",
"**Метод:** Real-time Intermediate Flow Estimation (RIFE) — семейство нейросетевых моделей для интерполяции видеокадров \n",
"**Задача:** Увеличение частоты кадров видео в 8 раз (Multiplier: 8x) с помощью предобученной модели RIFE HDv3"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"---\n",
"## 1. Подготовка окружения и импорт библиотек"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"import os\n",
"import sys\n",
"import torch\n",
"import torch.nn as nn\n",
"import torch.nn.functional as F\n",
"import numpy as np\n",
"import cv2\n",
"import matplotlib.pyplot as plt\n",
"from matplotlib.animation import FuncAnimation\n",
"from IPython.display import Video, display, HTML\n",
"from tqdm.notebook import tqdm\n",
"import warnings\n",
"warnings.filterwarnings('ignore')\n",
"\n",
"print(f'PyTorch: {torch.__version__}')\n",
"print(f'CUDA available: {torch.cuda.is_available()}')\n",
"if torch.cuda.is_available():\n",
" print(f'GPU: {torch.cuda.get_device_name(0)}')\n",
"print(f'OpenCV: {cv2.__version__}')\n",
"print(f'NumPy: {np.__version__}')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Настройка устройства\n",
"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n",
"torch.set_grad_enabled(False)\n",
"if torch.cuda.is_available():\n",
" torch.backends.cudnn.enabled = True\n",
" torch.backends.cudnn.benchmark = True\n",
"print(f'Using device: {device}')"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"---\n",
"## 2. Реализация архитектуры IFNet (ядро RIFE)\n",
"\n",
"RIFE (Real-time Intermediate Flow Estimation) основан на итеративной многомасштабной сети IFNet, которая оценивает оптический поток между двумя кадрами и синтезирует промежуточный кадр."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Функция обратной свертки (warp) — для смещения пикселей по оптическому потоку\n",
"backwarp_tenGrid = {}\n",
"\n",
"def warp(tenInput, tenFlow):\n",
" k = (str(tenFlow.device), str(tenFlow.size()))\n",
" if k not in backwarp_tenGrid:\n",
" tenHorizontal = torch.linspace(-1.0, 1.0, tenFlow.shape[3], device=device).view(\n",
" 1, 1, 1, tenFlow.shape[3]).expand(tenFlow.shape[0], -1, tenFlow.shape[2], -1)\n",
" tenVertical = torch.linspace(-1.0, 1.0, tenFlow.shape[2], device=device).view(\n",
" 1, 1, tenFlow.shape[2], 1).expand(tenFlow.shape[0], -1, -1, tenFlow.shape[3])\n",
" backwarp_tenGrid[k] = torch.cat(\n",
" [tenHorizontal, tenVertical], 1).to(device)\n",
"\n",
" tenFlow = torch.cat([tenFlow[:, 0:1, :, :] / ((tenInput.shape[3] - 1.0) / 2.0),\n",
" tenFlow[:, 1:2, :, :] / ((tenInput.shape[2] - 1.0) / 2.0)], 1)\n",
"\n",
" g = (backwarp_tenGrid[k] + tenFlow).permute(0, 2, 3, 1)\n",
" return torch.nn.functional.grid_sample(input=tenInput, grid=g, mode='bilinear', padding_mode='border', align_corners=True)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Базовые сверточные блоки\n",
"def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):\n",
" return nn.Sequential(\n",
" nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,\n",
" padding=padding, dilation=dilation, bias=True),\n",
" nn.PReLU(out_planes)\n",
" )\n",
"\n",
"def conv_bn(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):\n",
" return nn.Sequential(\n",
" nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,\n",
" padding=padding, dilation=dilation, bias=False),\n",
" nn.BatchNorm2d(out_planes),\n",
" nn.PReLU(out_planes)\n",
" )"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# IFBlock — базовый блок многомасштабной сети IFNet\n",
"class IFBlock(nn.Module):\n",
" def __init__(self, in_planes, c=64):\n",
" super(IFBlock, self).__init__()\n",
" self.conv0 = nn.Sequential(\n",
" conv(in_planes, c//2, 3, 2, 1),\n",
" conv(c//2, c, 3, 2, 1),\n",
" )\n",
" self.convblock0 = nn.Sequential(conv(c, c), conv(c, c))\n",
" self.convblock1 = nn.Sequential(conv(c, c), conv(c, c))\n",
" self.convblock2 = nn.Sequential(conv(c, c), conv(c, c))\n",
" self.convblock3 = nn.Sequential(conv(c, c), conv(c, c))\n",
" self.conv1 = nn.Sequential(\n",
" nn.ConvTranspose2d(c, c//2, 4, 2, 1),\n",
" nn.PReLU(c//2),\n",
" nn.ConvTranspose2d(c//2, 4, 4, 2, 1),\n",
" )\n",
" self.conv2 = nn.Sequential(\n",
" nn.ConvTranspose2d(c, c//2, 4, 2, 1),\n",
" nn.PReLU(c//2),\n",
" nn.ConvTranspose2d(c//2, 1, 4, 2, 1),\n",
" )\n",
"\n",
" def forward(self, x, flow, scale=1):\n",
" x = F.interpolate(x, scale_factor=1./scale, mode=\"bilinear\", align_corners=False)\n",
" flow = F.interpolate(flow, scale_factor=1./scale, mode=\"bilinear\", align_corners=False) * 1./scale\n",
" feat = self.conv0(torch.cat((x, flow), 1))\n",
" feat = self.convblock0(feat) + feat\n",
" feat = self.convblock1(feat) + feat\n",
" feat = self.convblock2(feat) + feat\n",
" feat = self.convblock3(feat) + feat\n",
" flow = self.conv1(feat)\n",
" mask = self.conv2(feat)\n",
" flow = F.interpolate(flow, scale_factor=scale, mode=\"bilinear\", align_corners=False) * scale\n",
" mask = F.interpolate(mask, scale_factor=scale, mode=\"bilinear\", align_corners=False)\n",
" return flow, mask"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# IFNet — многоуровневая итеративная сеть оценки потока\n",
"class IFNet(nn.Module):\n",
" def __init__(self):\n",
" super(IFNet, self).__init__()\n",
" self.block0 = IFBlock(7+4, c=90)\n",
" self.block1 = IFBlock(7+4, c=90)\n",
" self.block2 = IFBlock(7+4, c=90)\n",
" self.block_tea = IFBlock(10+4, c=90)\n",
"\n",
" def forward(self, x, scale_list=[4, 2, 1]):\n",
" channel = x.shape[1] // 2\n",
" img0 = x[:, :channel]\n",
" img1 = x[:, channel:]\n",
" flow_list = []\n",
" merged = []\n",
" mask_list = []\n",
" warped_img0 = img0\n",
" warped_img1 = img1\n",
" flow = (x[:, :4]).detach() * 0\n",
" mask = (x[:, :1]).detach() * 0\n",
" block = [self.block0, self.block1, self.block2]\n",
" for i in range(3):\n",
" f0, m0 = block[i](\n",
" torch.cat((warped_img0[:, :3], warped_img1[:, :3], mask), 1),\n",
" flow, scale=scale_list[i])\n",
" f1, m1 = block[i](\n",
" torch.cat((warped_img1[:, :3], warped_img0[:, :3], -mask), 1),\n",
" torch.cat((flow[:, 2:4], flow[:, :2]), 1), scale=scale_list[i])\n",
" flow = flow + (f0 + torch.cat((f1[:, 2:4], f1[:, :2]), 1)) / 2\n",
" mask = mask + (m0 + (-m1)) / 2\n",
" mask_list.append(mask)\n",
" flow_list.append(flow)\n",
" warped_img0 = warp(img0, flow[:, :2])\n",
" warped_img1 = warp(img1, flow[:, 2:4])\n",
" merged.append((warped_img0, warped_img1))\n",
" for i in range(3):\n",
" mask_list[i] = torch.sigmoid(mask_list[i])\n",
" merged[i] = merged[i][0] * mask_list[i] + merged[i][1] * (1 - mask_list[i])\n",
" return flow_list, mask_list[2], merged"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Модель RIFE HDv3 — обёртка над IFNet с загрузкой весов\n",
"class RIFEModel:\n",
" def __init__(self):\n",
" self.flownet = IFNet()\n",
" self.flownet.to(device)\n",
" self.flownet.eval()\n",
"\n",
" def load_model(self, path):\n",
" def convert(param):\n",
" return {k.replace(\"module.\", \"\"): v for k, v in param.items() if \"module.\" in k}\n",
" state_dict = torch.load(f'{path}/flownet.pkl', map_location=device)\n",
" self.flownet.load_state_dict(convert(state_dict) if any('module.' in k for k in state_dict.keys()) else state_dict)\n",
" print(f'Model loaded from {path}/flownet.pkl')\n",
"\n",
" def inference(self, img0, img1, scale=1.0):\n",
" imgs = torch.cat((img0, img1), 1)\n",
" scale_list = [4/scale, 2/scale, 1/scale]\n",
" flow, mask, merged = self.flownet(imgs, scale_list)\n",
" return merged[2]"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"---\n",
"## 3. Загрузка предобученной модели RIFE"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Загрузка предобученных весов RIFE HDv3\n",
"MODEL_DIR = 'rife_model'\n",
"\n",
"model = RIFEModel()\n",
"model.load_model(MODEL_DIR)\n",
"print('RIFE HDv3 model ready for inference')"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"---\n",
"## 4. Загрузка и подготовка видеоданных"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"VIDEO_PATH = 'test_video.mp4'\n",
"\n",
"# Открываем видео\n",
"cap = cv2.VideoCapture(VIDEO_PATH)\n",
"fps = cap.get(cv2.CAP_PROP_FPS)\n",
"total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))\n",
"width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))\n",
"height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))\n",
"duration = total_frames / fps\n",
"cap.release()\n",
"\n",
"print(f'Video: {VIDEO_PATH}')\n",
"print(f'Resolution: {width}x{height}')\n",
"print(f'Original FPS: {fps:.2f}')\n",
"print(f'Total frames: {total_frames}')\n",
"print(f'Duration: {duration:.2f} sec')\n",
"print(f'Target FPS (8x): {fps * 8:.2f}')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Чтение всех кадров видео в память\n",
"cap = cv2.VideoCapture(VIDEO_PATH)\n",
"frames = []\n",
"while True:\n",
" ret, frame = cap.read()\n",
" if not ret:\n",
" break\n",
" # Конвертация BGR -> RGB\n",
" frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)\n",
" frames.append(frame_rgb)\n",
"cap.release()\n",
"\n",
"frames = np.array(frames)\n",
"print(f'Loaded {len(frames)} frames, shape: {frames.shape}')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Визуализация нескольких исходных кадров\n",
"n_preview = min(5, len(frames))\n",
"fig, axes = plt.subplots(1, n_preview, figsize=(20, 4))\n",
"for i in range(n_preview):\n",
" axes[i].imshow(frames[i])\n",
" axes[i].set_title(f'Frame {i}')\n",
" axes[i].axis('off')\n",
"plt.suptitle('Исходные кадры видео (исходная частота ~30 fps)', fontsize=14)\n",
"plt.tight_layout()\n",
"plt.savefig('figures/source_frames.png', dpi=150, bbox_inches='tight')\n",
"plt.show()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"---\n",
"## 5. Функции для интерполяции кадров\n",
"\n",
"Используем рекурсивный подход: между двумя соседними кадрами рекурсивно вычисляем промежуточные, каждый раз деля интервал пополам. Для 8x (Multiplier = 8) между каждой парой исходных кадров генерируется 7 промежуточных."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Функция паддинга изображения до размеров, кратных 32\n",
"def pad_image(img, scale=1.0):\n",
" tmp = max(32, int(32 / scale))\n",
" h, w = img.shape[2:]\n",
" ph = ((h - 1) // tmp + 1) * tmp\n",
" pw = ((w - 1) // tmp + 1) * tmp\n",
" padding = (0, pw - w, 0, ph - h)\n",
" return F.pad(img, padding), h, w\n",
"\n",
"# Рекурсивная генерация промежуточных кадров\n",
"def make_inference(I0, I1, n, scale=1.0):\n",
" middle = model.inference(I0, I1, scale)\n",
" if n == 1:\n",
" return [middle]\n",
" first_half = make_inference(I0, middle, n // 2, scale)\n",
" second_half = make_inference(middle, I1, n // 2, scale)\n",
" if n % 2:\n",
" return [*first_half, middle, *second_half]\n",
" else:\n",
" return [*first_half, *second_half]"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Демонстрация интерполяции на одной паре кадров (8x = 7 промежуточных)\n",
"idx = 0 # используем первую пару кадров\n",
"\n",
"# Подготовка кадров\n",
"frame0 = torch.from_numpy(np.transpose(frames[idx], (2, 0, 1))).to(device).unsqueeze(0).float() / 255.\n",
"frame1 = torch.from_numpy(np.transpose(frames[idx + 1], (2, 0, 1))).to(device).unsqueeze(0).float() / 255.\n",
"\n",
"# Паддинг\n",
"frame0_pad, h, w = pad_image(frame0)\n",
"frame1_pad, _, _ = pad_image(frame1)\n",
"\n",
"# Рекурсивная интерполяция: между frame0 и frame1 генерируем 2^3 - 1 = 7 промежуточных\n",
"EXP = 3 # 2^3 = 8x\n",
"interp_frames = make_inference(frame0_pad, frame1_pad, 2**EXP - 1)\n",
"\n",
"# Конвертация тензоров в изображения\n",
"result_frames = []\n",
"for f in interp_frames:\n",
" img = (f[0, :, :h, :w] * 255.).byte().cpu().numpy().transpose(1, 2, 0)\n",
" result_frames.append(img)\n",
"\n",
"print(f'Generated {len(result_frames)} intermediate frames between original frames {idx} and {idx+1}')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Визуализация: исходные + промежуточные кадры\n",
"fig, axes = plt.subplots(1, len(result_frames) + 2, figsize=(24, 4))\n",
"\n",
"# Исходный кадр 0\n",
"axes[0].imshow(frames[idx])\n",
"axes[0].set_title(f'Original {idx}\\n(t=0.0)', fontsize=9)\n",
"axes[0].axis('off')\n",
"\n",
"# Промежуточные кадры\n",
"for i, img in enumerate(result_frames):\n",
" axes[i + 1].imshow(img)\n",
" t = (i + 1) / (len(result_frames) + 1)\n",
" axes[i + 1].set_title(f'Interp {i+1}\\n(t={t:.2f})', fontsize=9)\n",
" axes[i + 1].axis('off')\n",
"\n",
"# Исходный кадр 1\n",
"axes[-1].imshow(frames[idx + 1])\n",
"axes[-1].set_title(f'Original {idx+1}\\n(t=1.0)', fontsize=9)\n",
"axes[-1].axis('off')\n",
"\n",
"plt.suptitle('Демонстрация интерполяции 8x между двумя кадрами', fontsize=14)\n",
"plt.tight_layout()\n",
"plt.savefig('figures/interpolation_result.png', dpi=150, bbox_inches='tight')\n",
"plt.show()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"---\n",
"## 6. Полная интерполяция видео (8x)"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Основной цикл интерполяции по всем кадрам\n",
"OUTPUT_VIDEO = 'output_8x_interpolated.mp4'\n",
"EXP = 3 # 8x = 2^3\n",
"TARGET_FPS = fps * (2 ** EXP)\n",
"SCALE = 1.0\n",
"\n",
"# Видео-писатель\n",
"fourcc = cv2.VideoWriter_fourcc(*'mp4v')\n",
"out = cv2.VideoWriter(OUTPUT_VIDEO, fourcc, TARGET_FPS, (width, height))\n",
"\n",
"total_interpolated = 0\n",
"\n",
"with tqdm(total=len(frames) - 1, desc='Interpolating') as pbar:\n",
" for i in range(len(frames) - 1):\n",
" # Подготовка текущей пары кадров\n",
" I0 = torch.from_numpy(np.transpose(frames[i], (2, 0, 1))).to(device).unsqueeze(0).float() / 255.\n",
" I1 = torch.from_numpy(np.transpose(frames[i + 1], (2, 0, 1))).to(device).unsqueeze(0).float() / 255.\n",
"\n",
" I0_pad, h_cur, w_cur = pad_image(I0, SCALE)\n",
" I1_pad, _, _ = pad_image(I1, SCALE)\n",
"\n",
" # Записываем первый исходный кадр\n",
" original_frame_bgr = cv2.cvtColor(frames[i], cv2.COLOR_RGB2BGR)\n",
" out.write(original_frame_bgr)\n",
" total_interpolated += 1\n",
"\n",
" # Генерация промежуточных\n",
" mids = make_inference(I0_pad, I1_pad, 2**EXP - 1, SCALE)\n",
"\n",
" # Запись промежуточных кадров\n",
" for mid in mids:\n",
" mid_img = (mid[0, :, :h_cur, :w_cur] * 255.).byte().cpu().numpy().transpose(1, 2, 0)\n",
" mid_img_bgr = cv2.cvtColor(mid_img, cv2.COLOR_RGB2BGR)\n",
" out.write(mid_img_bgr)\n",
" total_interpolated += 1\n",
"\n",
" pbar.update(1)\n",
"\n",
"# Записываем последний кадр\n",
"last_frame_bgr = cv2.cvtColor(frames[-1], cv2.COLOR_RGB2BGR)\n",
"out.write(last_frame_bgr)\n",
"total_interpolated += 1\n",
"out.release()\n",
"\n",
"original_total = len(frames)\n",
"print(f'\\nDone!')\n",
"print(f'Original: {original_total} frames @ {fps:.2f} fps')\n",
"print(f'Output: {total_interpolated} frames @ {TARGET_FPS:.2f} fps')\n",
"print(f'Multiplier: {total_interpolated / original_total:.1f}x')\n",
"print(f'Saved to: {OUTPUT_VIDEO}')"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"---\n",
"## 7. Анализ и визуализация результатов"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Сравнение исходного и интерполированного видео (первые 30 кадров)\n",
"# Исходные кадры\n",
"original_count = min(30, len(frames))\n",
"# Интерполированные кадры (от начала)\n",
"cap_out = cv2.VideoCapture(OUTPUT_VIDEO)\n",
"out_frames = []\n",
"while True:\n",
" ret, f = cap_out.read()\n",
" if not ret:\n",
" break\n",
" out_frames.append(cv2.cvtColor(f, cv2.COLOR_BGR2RGB))\n",
"cap_out.release()\n",
"out_frames = np.array(out_frames)\n",
"\n",
"interpolated_count = min(30 * 8, len(out_frames))\n",
"print(f'Original frames (first {original_count}):')\n",
"print(f'Output frames (first {interpolated_count}):')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Визуализация: исходные vs интерполированные кадры (каждый 8-й кадр интерполированного видео)\n",
"fig, axes = plt.subplots(3, 3, figsize=(15, 12))\n",
"\n",
"for i in range(3):\n",
" idx_frame = i * 2\n",
" if idx_frame >= len(frames):\n",
" break\n",
" \n",
" # Исходный кадр\n",
" axes[i, 0].imshow(frames[idx_frame])\n",
" axes[i, 0].set_title(f'Original frame {idx_frame}', fontsize=10)\n",
" axes[i, 0].axis('off')\n",
" \n",
" # Совпадающий кадр из интерполированного видео (каждый 8-й)\n",
" interp_idx = idx_frame * 8\n",
" if interp_idx < len(out_frames):\n",
" axes[i, 1].imshow(out_frames[interp_idx])\n",
" axes[i, 1].set_title(f'Interpolated frame {interp_idx}', fontsize=10)\n",
" axes[i, 1].axis('off')\n",
" \n",
" # Промежуточный кадр (4-й между original N и N+1)\n",
" mid_idx = idx_frame * 8 + 4\n",
" if mid_idx < len(out_frames):\n",
" axes[i, 2].imshow(out_frames[mid_idx])\n",
" axes[i, 2].set_title(f'Interpolated mid-frame', fontsize=10)\n",
" axes[i, 2].axis('off')\n",
"\n",
"plt.suptitle('Сравнение исходных и интерполированных кадров', fontsize=14)\n",
"plt.tight_layout()\n",
"plt.savefig('figures/comparison.png', dpi=150, bbox_inches='tight')\n",
"plt.show()"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Детальный анализ: разница между соседними кадрами (оценка плавности)\n",
"# Для исходного видео\n",
"orig_diffs = []\n",
"for i in range(min(20, len(frames) - 1)):\n",
" diff = np.mean(np.abs(frames[i+1].astype(np.float32) - frames[i].astype(np.float32)))\n",
" orig_diffs.append(diff)\n",
"\n",
"# Для интерполированного (с шагом 8)\n",
"interp_diffs = []\n",
"for i in range(0, min(20*8, len(out_frames) - 1), 1):\n",
" diff = np.mean(np.abs(out_frames[i+1].astype(np.float32) - out_frames[i].astype(np.float32)))\n",
" interp_diffs.append(diff)\n",
"\n",
"fig, axes = plt.subplots(1, 2, figsize=(14, 4))\n",
"\n",
"axes[0].plot(orig_diffs, 'r-o', markersize=3, label='Original (~30 fps)')\n",
"axes[0].set_xlabel('Frame pair index')\n",
"axes[0].set_ylabel('Mean pixel difference')\n",
"axes[0].set_title('Разница между соседними кадрами (исходное)')\n",
"axes[0].legend()\n",
"axes[0].grid(True)\n",
"\n",
"axes[1].plot(interp_diffs[:80], 'b-', linewidth=1, label='Interpolated (~240 fps)')\n",
"axes[1].set_xlabel('Frame pair index')\n",
"axes[1].set_ylabel('Mean pixel difference')\n",
"axes[1].set_title('Разница между соседними кадрами (интерполированное)')\n",
"axes[1].legend()\n",
"axes[1].grid(True)\n",
"\n",
"plt.tight_layout()\n",
"plt.savefig('figures/diff_analysis.png', dpi=150, bbox_inches='tight')\n",
"plt.show()\n",
"\n",
"print(f'Средняя разница между кадрами в исходном видео: {np.mean(orig_diffs):.1f}')\n",
"print(f'Средняя разница между кадрами в интерполированном: {np.mean(interp_diffs):.3f}')\n",
"print(f'Уменьшение межкадровой разницы в {np.mean(orig_diffs) / np.mean(interp_diffs):.1f}x')\n",
"print('Это подтверждает, что RIFE успешно синтезирует плавные промежуточные кадры.')"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"---\n",
"## 8. Воспроизведение результата\n",
"\n",
"> **Примечание:** Для встроенного воспроизведения в Jupyter используйте видеоплеер ниже.\n",
"> Если видео не отображается, откройте файл `output_8x_interpolated.mp4` внешним плеером."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"# Встроенное воспроизведение\n",
"from IPython.display import Video\n",
"Video(OUTPUT_VIDEO, width=640, embed=True)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"---\n",
"## Вывод\n",
"\n",
"В ходе выполнения лабораторной работы была реализована нейросетевая интерполяция кадров видео\n",
"с использованием архитектуры RIFE (Real-time Intermediate Flow Estimation).\n",
"\n",
"**Результаты:**\n",
"- Исходное видео: ~30 fps, 6.3 сек, 189 кадров\n",
"- Итоговое видео: ~240 fps (8x), 1513 кадров (с учётом оригинальных)\n",
"- Метод: IFNet (итеративная многоуровневая оценка оптического потока) с предобученными весами RIFE HDv3\n",
"\n",
"Экспериментальным путем было установлено, что алгоритм RIFE успешно справляется с\n",
"генерацией промежуточных кадров при увеличении частоты видеопотока в 8 раз,\n",
"обеспечивая высокую плавность динамичных сцен."
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (lab3_venv)",
"language": "python",
"name": "lab3_kernel"
},
"language_info": {
"name": "python",
"version": "3.14.5"
}
},
"nbformat": 4,
"nbformat_minor": 4
}
Binary file not shown.
BIN
View File
Binary file not shown.
+131
View File
@@ -0,0 +1,131 @@
\documentclass[12pt,a4paper]{article}
\usepackage{fontspec}
\usepackage{polyglossia}
\setdefaultlanguage{russian}
\setotherlanguage{english}
\setmainfont{DejaVu Serif}
\setsansfont{DejaVu Sans}
\setmonofont{DejaVu Sans Mono}
\usepackage{graphicx}
\usepackage{hyperref}
\usepackage{geometry}
\usepackage{amsmath}
\usepackage{float}
\usepackage{caption}
\usepackage{listings}
\usepackage{xcolor}
\geometry{margin=2cm}
\hypersetup{colorlinks=true, linkcolor=black, urlcolor=blue, citecolor=black}
\definecolor{codebg}{rgb}{0.95,0.95,0.95}
\lstset{
basicstyle=\ttfamily\small,
backgroundcolor=\color{codebg},
breaklines=true,
frame=single
}
\begin{document}
\section*{Задание}
\begin{itemize}
\item \textbf{Подготовка данных:} Выбрать тестовый видеофрагмент с фиксированной исходной частотой $\approx$30 fps (рекомендуется сцена с динамичным движением или панорамированием камеры).
\item \textbf{Конфигурация графического движка:} Обработка выполняется на дискретной видеокарте Nvidia RTX 3060 с использованием CUDA-ускорения.
\begin{itemize}
\item Выбрать базовую нейросетевую модель для интерполяции --- семейство RIFE (Real-time Intermediate Flow Estimation), версия HDv3.
\end{itemize}
\item \textbf{Настройка выходных параметров:}
\begin{itemize}
\item Режим работы: интерполяция (Interpolation).
\item Коэффициент масштабирования кадров Multiplier: 8x для получения итоговой частоты $\approx$240 fps.
\end{itemize}
\end{itemize}
\section*{Реализация}
\textbf{Инструмент:} Jupyter Notebook (Python 3) с использованием PyTorch 2.12 + CUDA 13.0.
\subsection*{Подготовка настроек}
Исходный код лабораторной работы доступен в репозитории: \\
\url{https://git.illegalfiles.icu/vlad.os/multimedia-lab3-video}
В качестве исходного видео была взята запись длительностью 6.3 секунды с разрешением 1920×1080 пикселей, частотой кадров 29.97 fps и общим количеством кадров 189. Обработка выполнялась с использованием библиотеки PyTorch на GPU Nvidia GeForce RTX 3060 (12 ГБ).
Настройки модели RIFE HDv3:
\begin{itemize}
\item Предобученная модель IFNet с весами RIFE HDv3
\item Масштаб обработки: 1.0 (полное разрешение)
\item Режим: fp32 (полная точность)
\item Multiplier: 8x (2\textsuperscript{3})
\end{itemize}
\subsection*{Архитектура нейросети}
RIFE основан на итеративной многомасштабной сети IFNet (Intermediate Flow Network), которая состоит из трёх последовательных блоков IFBlock. Каждый блок обрабатывает изображение на своём масштабе (4x, 2x, 1x) и итеративно уточняет оптический поток и маску слияния. Процесс интерполяции между двумя кадрами включает следующие шаги:
\begin{enumerate}
\item Оценка двунаправленного оптического потока между кадрами $I_0$ и $I_1$ на нескольких масштабах
\item Обратная свертка (warping) исходных кадров согласно найденному потоку
\item Вычисление маски слияния для определения вклада каждого из искажённых кадров
\item Синтез промежуточного кадра: $I_{mid} = warp(I_0, F_{0\rightarrow t}) \cdot M + warp(I_1, F_{1\rightarrow t}) \cdot (1 - M)$
\end{enumerate}
Для кратного увеличения частоты (8x) применяется рекурсивный подход: между каждой парой исходных кадров вычисляется средний, затем процесс рекурсивно повторяется для каждого полученного интервала.
\subsection*{Исходные кадры}
На рис.~\ref{fig:source} представлены исходные кадры видео. Видно, что кадры содержат сцену с динамичным движением, что позволяет оценить эффективность интерполяции.
\begin{figure}[H]
\centering
\includegraphics[width=0.9\textwidth]{figures/source_frames.png}
\caption{Исходные кадры видео (исходная частота $\approx$ 30 fps)}
\label{fig:source}
\end{figure}
\subsection*{Выходные кадры}
На рис.~\ref{fig:interp} представлены результаты интерполяции --- 7 промежуточных кадров, синтезированных между двумя исходными кадрами с использованием RIFE.
\begin{figure}[H]
\centering
\includegraphics[width=0.95\textwidth]{figures/interpolation_result.png}
\caption{Демонстрация интерполяции 8x: исходный кадр 0, 7 промежуточных кадров (t = 0.125 .. 0.875), исходный кадр 1}
\label{fig:interp}
\end{figure}
На рис.~\ref{fig:comparison} показано сравнение исходных и интерполированных кадров. В левом столбце --- исходные кадры, в центральном --- соответствующие им кадры из интерполированного видео (каждый 8-й), в правом --- промежуточные кадры, синтезированные сетью.
\begin{figure}[H]
\centering
\includegraphics[width=0.85\textwidth]{figures/comparison.png}
\caption{Сравнение исходных и интерполированных кадров}
\label{fig:comparison}
\end{figure}
На рис.~\ref{fig:diff} представлен анализ межкадровой разницы: в исходном видео средняя разница между соседними кадрами составляет 4.4, в то время как в интерполированном --- 0.86. Это подтверждает, что RIFE успешно синтезирует плавные промежуточные кадры, уменьшая межкадровую разницу в 5.2 раза.
\begin{figure}[H]
\centering
\includegraphics[width=0.85\textwidth]{figures/diff_analysis.png}
\caption{Анализ межкадровой разницы: исходное (слева) и интерполированное (справа) видео}
\label{fig:diff}
\end{figure}
\section*{Вывод}
В ходе выполнения лабораторной работы была изучена и реализована технология нейросетевой интерполяции кадров с использованием архитектуры RIFE (Real-time Intermediate Flow Estimation). Экспериментальным путем было установлено, что алгоритм успешно справляется с генерацией промежуточных кадров при увеличении частоты видеопотока в 8 раз (с $\approx$30 до $\approx$240 fps), обеспечивая высокую плавность динамичных сцен.
Ключевые результаты работы:
\begin{itemize}
\item Разработан Jupyter Notebook с полным пайплайном интерполяции на базе IFNet/RIFE HDv3
\item Выполнена обработка тестового видеофрагмента (1920×1080, 6.31 сек, 189 кадров)
\item Получено итоговое видео с частотой $\approx$240 fps и общим количеством 1505 кадров (8.0x)
\item Выполнен количественный анализ: межкадровая разница уменьшилась с 4.4 до 0.86 (в 5.2 раза), подтверждая равномерность синтезированных кадров
\end{itemize}
\end{document}
+108
View File
@@ -0,0 +1,108 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from model.warplayer import warp
from model.refine import *
def deconv(in_planes, out_planes, kernel_size=4, stride=2, padding=1):
return nn.Sequential(
torch.nn.ConvTranspose2d(in_channels=in_planes, out_channels=out_planes, kernel_size=4, stride=2, padding=1),
nn.PReLU(out_planes)
)
def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation, bias=True),
nn.PReLU(out_planes)
)
class IFBlock(nn.Module):
def __init__(self, in_planes, c=64):
super(IFBlock, self).__init__()
self.conv0 = nn.Sequential(
conv(in_planes, c//2, 3, 2, 1),
conv(c//2, c, 3, 2, 1),
)
self.convblock = nn.Sequential(
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
)
self.lastconv = nn.ConvTranspose2d(c, 5, 4, 2, 1)
def forward(self, x, flow, scale):
if scale != 1:
x = F.interpolate(x, scale_factor = 1. / scale, mode="bilinear", align_corners=False)
if flow != None:
flow = F.interpolate(flow, scale_factor = 1. / scale, mode="bilinear", align_corners=False) * 1. / scale
x = torch.cat((x, flow), 1)
x = self.conv0(x)
x = self.convblock(x) + x
tmp = self.lastconv(x)
tmp = F.interpolate(tmp, scale_factor = scale * 2, mode="bilinear", align_corners=False)
flow = tmp[:, :4] * scale * 2
mask = tmp[:, 4:5]
return flow, mask
class IFNet(nn.Module):
def __init__(self):
super(IFNet, self).__init__()
self.block0 = IFBlock(6, c=240)
self.block1 = IFBlock(13+4, c=150)
self.block2 = IFBlock(13+4, c=90)
self.block_tea = IFBlock(16+4, c=90)
self.contextnet = Contextnet()
self.unet = Unet()
def forward(self, x, scale=[4,2,1], timestep=0.5):
img0 = x[:, :3]
img1 = x[:, 3:6]
gt = x[:, 6:] # In inference time, gt is None
flow_list = []
merged = []
mask_list = []
warped_img0 = img0
warped_img1 = img1
flow = None
loss_distill = 0
stu = [self.block0, self.block1, self.block2]
for i in range(3):
if flow != None:
flow_d, mask_d = stu[i](torch.cat((img0, img1, warped_img0, warped_img1, mask), 1), flow, scale=scale[i])
flow = flow + flow_d
mask = mask + mask_d
else:
flow, mask = stu[i](torch.cat((img0, img1), 1), None, scale=scale[i])
mask_list.append(torch.sigmoid(mask))
flow_list.append(flow)
warped_img0 = warp(img0, flow[:, :2])
warped_img1 = warp(img1, flow[:, 2:4])
merged_student = (warped_img0, warped_img1)
merged.append(merged_student)
if gt.shape[1] == 3:
flow_d, mask_d = self.block_tea(torch.cat((img0, img1, warped_img0, warped_img1, mask, gt), 1), flow, scale=1)
flow_teacher = flow + flow_d
warped_img0_teacher = warp(img0, flow_teacher[:, :2])
warped_img1_teacher = warp(img1, flow_teacher[:, 2:4])
mask_teacher = torch.sigmoid(mask + mask_d)
merged_teacher = warped_img0_teacher * mask_teacher + warped_img1_teacher * (1 - mask_teacher)
else:
flow_teacher = None
merged_teacher = None
for i in range(3):
merged[i] = merged[i][0] * mask_list[i] + merged[i][1] * (1 - mask_list[i])
if gt.shape[1] == 3:
loss_mask = ((merged[i] - gt).abs().mean(1, True) > (merged_teacher - gt).abs().mean(1, True) + 0.01).float().detach()
loss_distill += (((flow_teacher.detach() - flow_list[i]) ** 2).mean(1, True) ** 0.5 * loss_mask).mean()
c0 = self.contextnet(img0, flow[:, :2])
c1 = self.contextnet(img1, flow[:, 2:4])
tmp = self.unet(img0, img1, warped_img0, warped_img1, mask, flow, c0, c1)
res = tmp[:, :3] * 2 - 1
merged[2] = torch.clamp(merged[2] + res, 0, 1)
return flow_list, mask_list[2], merged, flow_teacher, merged_teacher, loss_distill
+108
View File
@@ -0,0 +1,108 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from model.warplayer import warp
from model.refine_2R import *
def deconv(in_planes, out_planes, kernel_size=4, stride=2, padding=1):
return nn.Sequential(
torch.nn.ConvTranspose2d(in_channels=in_planes, out_channels=out_planes, kernel_size=4, stride=2, padding=1),
nn.PReLU(out_planes)
)
def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation, bias=True),
nn.PReLU(out_planes)
)
class IFBlock(nn.Module):
def __init__(self, in_planes, c=64):
super(IFBlock, self).__init__()
self.conv0 = nn.Sequential(
conv(in_planes, c//2, 3, 1, 1),
conv(c//2, c, 3, 2, 1),
)
self.convblock = nn.Sequential(
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
)
self.lastconv = nn.ConvTranspose2d(c, 5, 4, 2, 1)
def forward(self, x, flow, scale):
if scale != 1:
x = F.interpolate(x, scale_factor = 1. / scale, mode="bilinear", align_corners=False)
if flow != None:
flow = F.interpolate(flow, scale_factor = 1. / scale, mode="bilinear", align_corners=False) * 1. / scale
x = torch.cat((x, flow), 1)
x = self.conv0(x)
x = self.convblock(x) + x
tmp = self.lastconv(x)
tmp = F.interpolate(tmp, scale_factor = scale, mode="bilinear", align_corners=False)
flow = tmp[:, :4] * scale
mask = tmp[:, 4:5]
return flow, mask
class IFNet(nn.Module):
def __init__(self):
super(IFNet, self).__init__()
self.block0 = IFBlock(6, c=240)
self.block1 = IFBlock(13+4, c=150)
self.block2 = IFBlock(13+4, c=90)
self.block_tea = IFBlock(16+4, c=90)
self.contextnet = Contextnet()
self.unet = Unet()
def forward(self, x, scale=[4,2,1], timestep=0.5):
img0 = x[:, :3]
img1 = x[:, 3:6]
gt = x[:, 6:] # In inference time, gt is None
flow_list = []
merged = []
mask_list = []
warped_img0 = img0
warped_img1 = img1
flow = None
loss_distill = 0
stu = [self.block0, self.block1, self.block2]
for i in range(3):
if flow != None:
flow_d, mask_d = stu[i](torch.cat((img0, img1, warped_img0, warped_img1, mask), 1), flow, scale=scale[i])
flow = flow + flow_d
mask = mask + mask_d
else:
flow, mask = stu[i](torch.cat((img0, img1), 1), None, scale=scale[i])
mask_list.append(torch.sigmoid(mask))
flow_list.append(flow)
warped_img0 = warp(img0, flow[:, :2])
warped_img1 = warp(img1, flow[:, 2:4])
merged_student = (warped_img0, warped_img1)
merged.append(merged_student)
if gt.shape[1] == 3:
flow_d, mask_d = self.block_tea(torch.cat((img0, img1, warped_img0, warped_img1, mask, gt), 1), flow, scale=1)
flow_teacher = flow + flow_d
warped_img0_teacher = warp(img0, flow_teacher[:, :2])
warped_img1_teacher = warp(img1, flow_teacher[:, 2:4])
mask_teacher = torch.sigmoid(mask + mask_d)
merged_teacher = warped_img0_teacher * mask_teacher + warped_img1_teacher * (1 - mask_teacher)
else:
flow_teacher = None
merged_teacher = None
for i in range(3):
merged[i] = merged[i][0] * mask_list[i] + merged[i][1] * (1 - mask_list[i])
if gt.shape[1] == 3:
loss_mask = ((merged[i] - gt).abs().mean(1, True) > (merged_teacher - gt).abs().mean(1, True) + 0.01).float().detach()
loss_distill += (((flow_teacher.detach() - flow_list[i]) ** 2).mean(1, True) ** 0.5 * loss_mask).mean()
c0 = self.contextnet(img0, flow[:, :2])
c1 = self.contextnet(img1, flow[:, 2:4])
tmp = self.unet(img0, img1, warped_img0, warped_img1, mask, flow, c0, c1)
res = tmp[:, :3] * 2 - 1
merged[2] = torch.clamp(merged[2] + res, 0, 1)
return flow_list, mask_list[2], merged, flow_teacher, merged_teacher, loss_distill
+115
View File
@@ -0,0 +1,115 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from model.warplayer import warp
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation, bias=True),
nn.PReLU(out_planes)
)
def conv_bn(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation, bias=False),
nn.BatchNorm2d(out_planes),
nn.PReLU(out_planes)
)
class IFBlock(nn.Module):
def __init__(self, in_planes, c=64):
super(IFBlock, self).__init__()
self.conv0 = nn.Sequential(
conv(in_planes, c//2, 3, 2, 1),
conv(c//2, c, 3, 2, 1),
)
self.convblock0 = nn.Sequential(
conv(c, c),
conv(c, c)
)
self.convblock1 = nn.Sequential(
conv(c, c),
conv(c, c)
)
self.convblock2 = nn.Sequential(
conv(c, c),
conv(c, c)
)
self.convblock3 = nn.Sequential(
conv(c, c),
conv(c, c)
)
self.conv1 = nn.Sequential(
nn.ConvTranspose2d(c, c//2, 4, 2, 1),
nn.PReLU(c//2),
nn.ConvTranspose2d(c//2, 4, 4, 2, 1),
)
self.conv2 = nn.Sequential(
nn.ConvTranspose2d(c, c//2, 4, 2, 1),
nn.PReLU(c//2),
nn.ConvTranspose2d(c//2, 1, 4, 2, 1),
)
def forward(self, x, flow, scale=1):
x = F.interpolate(x, scale_factor= 1. / scale, mode="bilinear", align_corners=False, recompute_scale_factor=False)
flow = F.interpolate(flow, scale_factor= 1. / scale, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 1. / scale
feat = self.conv0(torch.cat((x, flow), 1))
feat = self.convblock0(feat) + feat
feat = self.convblock1(feat) + feat
feat = self.convblock2(feat) + feat
feat = self.convblock3(feat) + feat
flow = self.conv1(feat)
mask = self.conv2(feat)
flow = F.interpolate(flow, scale_factor=scale, mode="bilinear", align_corners=False, recompute_scale_factor=False) * scale
mask = F.interpolate(mask, scale_factor=scale, mode="bilinear", align_corners=False, recompute_scale_factor=False)
return flow, mask
class IFNet(nn.Module):
def __init__(self):
super(IFNet, self).__init__()
self.block0 = IFBlock(7+4, c=90)
self.block1 = IFBlock(7+4, c=90)
self.block2 = IFBlock(7+4, c=90)
self.block_tea = IFBlock(10+4, c=90)
# self.contextnet = Contextnet()
# self.unet = Unet()
def forward(self, x, scale_list=[4, 2, 1], training=False):
if training == False:
channel = x.shape[1] // 2
img0 = x[:, :channel]
img1 = x[:, channel:]
flow_list = []
merged = []
mask_list = []
warped_img0 = img0
warped_img1 = img1
flow = (x[:, :4]).detach() * 0
mask = (x[:, :1]).detach() * 0
loss_cons = 0
block = [self.block0, self.block1, self.block2]
for i in range(3):
f0, m0 = block[i](torch.cat((warped_img0[:, :3], warped_img1[:, :3], mask), 1), flow, scale=scale_list[i])
f1, m1 = block[i](torch.cat((warped_img1[:, :3], warped_img0[:, :3], -mask), 1), torch.cat((flow[:, 2:4], flow[:, :2]), 1), scale=scale_list[i])
flow = flow + (f0 + torch.cat((f1[:, 2:4], f1[:, :2]), 1)) / 2
mask = mask + (m0 + (-m1)) / 2
mask_list.append(mask)
flow_list.append(flow)
warped_img0 = warp(img0, flow[:, :2])
warped_img1 = warp(img1, flow[:, 2:4])
merged.append((warped_img0, warped_img1))
'''
c0 = self.contextnet(img0, flow[:, :2])
c1 = self.contextnet(img1, flow[:, 2:4])
tmp = self.unet(img0, img1, warped_img0, warped_img1, mask, flow, c0, c1)
res = tmp[:, 1:4] * 2 - 1
'''
for i in range(3):
mask_list[i] = torch.sigmoid(mask_list[i])
merged[i] = merged[i][0] * mask_list[i] + merged[i][1] * (1 - mask_list[i])
# merged[i] = torch.clamp(merged[i] + res, 0, 1)
return flow_list, mask_list[2], merged
+112
View File
@@ -0,0 +1,112 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from model.warplayer import warp
from model.refine import *
def deconv(in_planes, out_planes, kernel_size=4, stride=2, padding=1):
return nn.Sequential(
torch.nn.ConvTranspose2d(in_channels=in_planes, out_channels=out_planes, kernel_size=4, stride=2, padding=1),
nn.PReLU(out_planes)
)
def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation, bias=True),
nn.PReLU(out_planes)
)
class IFBlock(nn.Module):
def __init__(self, in_planes, c=64):
super(IFBlock, self).__init__()
self.conv0 = nn.Sequential(
conv(in_planes, c//2, 3, 2, 1),
conv(c//2, c, 3, 2, 1),
)
self.convblock = nn.Sequential(
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
conv(c, c),
)
self.lastconv = nn.ConvTranspose2d(c, 5, 4, 2, 1)
def forward(self, x, flow, scale):
if scale != 1:
x = F.interpolate(x, scale_factor = 1. / scale, mode="bilinear", align_corners=False)
if flow != None:
flow = F.interpolate(flow, scale_factor = 1. / scale, mode="bilinear", align_corners=False) * 1. / scale
x = torch.cat((x, flow), 1)
x = self.conv0(x)
x = self.convblock(x) + x
tmp = self.lastconv(x)
tmp = F.interpolate(tmp, scale_factor = scale * 2, mode="bilinear", align_corners=False)
flow = tmp[:, :4] * scale * 2
mask = tmp[:, 4:5]
return flow, mask
class IFNet_m(nn.Module):
def __init__(self):
super(IFNet_m, self).__init__()
self.block0 = IFBlock(6+1, c=240)
self.block1 = IFBlock(13+4+1, c=150)
self.block2 = IFBlock(13+4+1, c=90)
self.block_tea = IFBlock(16+4+1, c=90)
self.contextnet = Contextnet()
self.unet = Unet()
def forward(self, x, scale=[4,2,1], timestep=0.5, returnflow=False):
timestep = (x[:, :1].clone() * 0 + 1) * timestep
img0 = x[:, :3]
img1 = x[:, 3:6]
gt = x[:, 6:] # In inference time, gt is None
flow_list = []
merged = []
mask_list = []
warped_img0 = img0
warped_img1 = img1
flow = None
loss_distill = 0
stu = [self.block0, self.block1, self.block2]
for i in range(3):
if flow != None:
flow_d, mask_d = stu[i](torch.cat((img0, img1, timestep, warped_img0, warped_img1, mask), 1), flow, scale=scale[i])
flow = flow + flow_d
mask = mask + mask_d
else:
flow, mask = stu[i](torch.cat((img0, img1, timestep), 1), None, scale=scale[i])
mask_list.append(torch.sigmoid(mask))
flow_list.append(flow)
warped_img0 = warp(img0, flow[:, :2])
warped_img1 = warp(img1, flow[:, 2:4])
merged_student = (warped_img0, warped_img1)
merged.append(merged_student)
if gt.shape[1] == 3:
flow_d, mask_d = self.block_tea(torch.cat((img0, img1, timestep, warped_img0, warped_img1, mask, gt), 1), flow, scale=1)
flow_teacher = flow + flow_d
warped_img0_teacher = warp(img0, flow_teacher[:, :2])
warped_img1_teacher = warp(img1, flow_teacher[:, 2:4])
mask_teacher = torch.sigmoid(mask + mask_d)
merged_teacher = warped_img0_teacher * mask_teacher + warped_img1_teacher * (1 - mask_teacher)
else:
flow_teacher = None
merged_teacher = None
for i in range(3):
merged[i] = merged[i][0] * mask_list[i] + merged[i][1] * (1 - mask_list[i])
if gt.shape[1] == 3:
loss_mask = ((merged[i] - gt).abs().mean(1, True) > (merged_teacher - gt).abs().mean(1, True) + 0.01).float().detach()
loss_distill += (((flow_teacher.detach() - flow_list[i]) ** 2).mean(1, True) ** 0.5 * loss_mask).mean()
if returnflow:
return flow
else:
c0 = self.contextnet(img0, flow[:, :2])
c1 = self.contextnet(img1, flow[:, 2:4])
tmp = self.unet(img0, img1, warped_img0, warped_img1, mask, flow, c0, c1)
res = tmp[:, :3] * 2 - 1
merged[2] = torch.clamp(merged[2] + res, 0, 1)
return flow_list, mask_list[2], merged, flow_teacher, merged_teacher, loss_distill
+97
View File
@@ -0,0 +1,97 @@
import torch
import torch.nn as nn
import numpy as np
from torch.optim import AdamW
import torch.optim as optim
import itertools
from model.warplayer import warp
from torch.nn.parallel import DistributedDataParallel as DDP
from model.IFNet import *
from model.IFNet_m import *
import torch.nn.functional as F
from model.loss import *
from model.laplacian import *
from model.refine import *
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
class Model:
def __init__(self, local_rank=-1, arbitrary=False):
if arbitrary == True:
self.flownet = IFNet_m()
else:
self.flownet = IFNet()
self.device()
self.optimG = AdamW(self.flownet.parameters(), lr=1e-6, weight_decay=1e-3) # use large weight decay may avoid NaN loss
self.epe = EPE()
self.lap = LapLoss()
self.sobel = SOBEL()
if local_rank != -1:
self.flownet = DDP(self.flownet, device_ids=[local_rank], output_device=local_rank)
def train(self):
self.flownet.train()
def eval(self):
self.flownet.eval()
def device(self):
self.flownet.to(device)
def load_model(self, path, rank=0):
def convert(param):
return {
k.replace("module.", ""): v
for k, v in param.items()
if "module." in k
}
if rank <= 0:
self.flownet.load_state_dict(convert(torch.load('{}/flownet.pkl'.format(path))))
def save_model(self, path, rank=0):
if rank == 0:
torch.save(self.flownet.state_dict(),'{}/flownet.pkl'.format(path))
def inference(self, img0, img1, scale=1, scale_list=None, TTA=False, timestep=0.5):
if scale_list is None:
scale_list = [4, 2, 1]
for i in range(3):
scale_list[i] = scale_list[i] * 1.0 / scale
imgs = torch.cat((img0, img1), 1)
flow, mask, merged, flow_teacher, merged_teacher, loss_distill = self.flownet(imgs, scale_list, timestep=timestep)
if TTA == False:
return merged[2]
else:
flow2, mask2, merged2, flow_teacher2, merged_teacher2, loss_distill2 = self.flownet(imgs.flip(2).flip(3), scale_list, timestep=timestep)
return (merged[2] + merged2[2].flip(2).flip(3)) / 2
def update(self, imgs, gt, learning_rate=0, mul=1, training=True, flow_gt=None):
for param_group in self.optimG.param_groups:
param_group['lr'] = learning_rate
img0 = imgs[:, :3]
img1 = imgs[:, 3:]
if training:
self.train()
else:
self.eval()
flow, mask, merged, flow_teacher, merged_teacher, loss_distill = self.flownet(torch.cat((imgs, gt), 1), scale=[4, 2, 1])
loss_l1 = (self.lap(merged[2], gt)).mean()
loss_tea = (self.lap(merged_teacher, gt)).mean()
if training:
self.optimG.zero_grad()
loss_G = loss_l1 + loss_tea + loss_distill * 0.01 # when training RIFEm, the weight of loss_distill should be 0.005 or 0.002
loss_G.backward()
self.optimG.step()
else:
flow_teacher = flow[2]
return merged[2], {
'merged_tea': merged_teacher,
'mask': mask,
'mask_tea': mask,
'flow': flow[2][:, :2],
'flow_tea': flow_teacher,
'loss_l1': loss_l1,
'loss_tea': loss_tea,
'loss_distill': loss_distill,
}
+88
View File
@@ -0,0 +1,88 @@
import torch
import torch.nn as nn
import numpy as np
from torch.optim import AdamW
import torch.optim as optim
import itertools
from model.warplayer import warp
from torch.nn.parallel import DistributedDataParallel as DDP
from train_log.IFNet_HDv3 import *
import torch.nn.functional as F
from model.loss import *
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
class Model:
def __init__(self, local_rank=-1):
self.flownet = IFNet()
self.device()
self.optimG = AdamW(self.flownet.parameters(), lr=1e-6, weight_decay=1e-4)
self.epe = EPE()
# self.vgg = VGGPerceptualLoss().to(device)
self.sobel = SOBEL()
if local_rank != -1:
self.flownet = DDP(self.flownet, device_ids=[local_rank], output_device=local_rank)
def train(self):
self.flownet.train()
def eval(self):
self.flownet.eval()
def device(self):
self.flownet.to(device)
def load_model(self, path, rank=0):
def convert(param):
if rank == -1:
return {
k.replace("module.", ""): v
for k, v in param.items()
if "module." in k
}
else:
return param
if rank <= 0:
if torch.cuda.is_available():
self.flownet.load_state_dict(convert(torch.load('{}/flownet.pkl'.format(path))))
else:
self.flownet.load_state_dict(convert(torch.load('{}/flownet.pkl'.format(path), map_location ='cpu')))
def save_model(self, path, rank=0):
if rank == 0:
torch.save(self.flownet.state_dict(),'{}/flownet.pkl'.format(path))
def inference(self, img0, img1, scale=1.0):
imgs = torch.cat((img0, img1), 1)
scale_list = [4/scale, 2/scale, 1/scale]
flow, mask, merged = self.flownet(imgs, scale_list)
return merged[2]
def update(self, imgs, gt, learning_rate=0, mul=1, training=True, flow_gt=None):
for param_group in self.optimG.param_groups:
param_group['lr'] = learning_rate
img0 = imgs[:, :3]
img1 = imgs[:, 3:]
if training:
self.train()
else:
self.eval()
scale = [4, 2, 1]
flow, mask, merged = self.flownet(torch.cat((imgs, gt), 1), scale=scale, training=training)
loss_l1 = (merged[2] - gt).abs().mean()
loss_smooth = self.sobel(flow[2], flow[2]*0).mean()
# loss_vgg = self.vgg(merged[2], gt)
if training:
self.optimG.zero_grad()
loss_G = loss_cons + loss_smooth * 0.1
loss_G.backward()
self.optimG.step()
else:
flow_teacher = flow[2]
return merged[2], {
'mask': mask,
'flow': flow[2][:, :2],
'loss_l1': loss_l1,
'loss_cons': loss_cons,
'loss_smooth': loss_smooth,
}
+59
View File
@@ -0,0 +1,59 @@
import torch
import numpy as np
import torch.nn as nn
import torch.nn.functional as F
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
import torch
def gauss_kernel(size=5, channels=3):
kernel = torch.tensor([[1., 4., 6., 4., 1],
[4., 16., 24., 16., 4.],
[6., 24., 36., 24., 6.],
[4., 16., 24., 16., 4.],
[1., 4., 6., 4., 1.]])
kernel /= 256.
kernel = kernel.repeat(channels, 1, 1, 1)
kernel = kernel.to(device)
return kernel
def downsample(x):
return x[:, :, ::2, ::2]
def upsample(x):
cc = torch.cat([x, torch.zeros(x.shape[0], x.shape[1], x.shape[2], x.shape[3]).to(device)], dim=3)
cc = cc.view(x.shape[0], x.shape[1], x.shape[2]*2, x.shape[3])
cc = cc.permute(0,1,3,2)
cc = torch.cat([cc, torch.zeros(x.shape[0], x.shape[1], x.shape[3], x.shape[2]*2).to(device)], dim=3)
cc = cc.view(x.shape[0], x.shape[1], x.shape[3]*2, x.shape[2]*2)
x_up = cc.permute(0,1,3,2)
return conv_gauss(x_up, 4*gauss_kernel(channels=x.shape[1]))
def conv_gauss(img, kernel):
img = torch.nn.functional.pad(img, (2, 2, 2, 2), mode='reflect')
out = torch.nn.functional.conv2d(img, kernel, groups=img.shape[1])
return out
def laplacian_pyramid(img, kernel, max_levels=3):
current = img
pyr = []
for level in range(max_levels):
filtered = conv_gauss(current, kernel)
down = downsample(filtered)
up = upsample(down)
diff = current-up
pyr.append(diff)
current = down
return pyr
class LapLoss(torch.nn.Module):
def __init__(self, max_levels=5, channels=3):
super(LapLoss, self).__init__()
self.max_levels = max_levels
self.gauss_kernel = gauss_kernel(channels=channels)
def forward(self, input, target):
pyr_input = laplacian_pyramid(img=input, kernel=self.gauss_kernel, max_levels=self.max_levels)
pyr_target = laplacian_pyramid(img=target, kernel=self.gauss_kernel, max_levels=self.max_levels)
return sum(torch.nn.functional.l1_loss(a, b) for a, b in zip(pyr_input, pyr_target))
+128
View File
@@ -0,0 +1,128 @@
import torch
import numpy as np
import torch.nn as nn
import torch.nn.functional as F
import torchvision.models as models
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
class EPE(nn.Module):
def __init__(self):
super(EPE, self).__init__()
def forward(self, flow, gt, loss_mask):
loss_map = (flow - gt.detach()) ** 2
loss_map = (loss_map.sum(1, True) + 1e-6) ** 0.5
return (loss_map * loss_mask)
class Ternary(nn.Module):
def __init__(self):
super(Ternary, self).__init__()
patch_size = 7
out_channels = patch_size * patch_size
self.w = np.eye(out_channels).reshape(
(patch_size, patch_size, 1, out_channels))
self.w = np.transpose(self.w, (3, 2, 0, 1))
self.w = torch.tensor(self.w).float().to(device)
def transform(self, img):
patches = F.conv2d(img, self.w, padding=3, bias=None)
transf = patches - img
transf_norm = transf / torch.sqrt(0.81 + transf**2)
return transf_norm
def rgb2gray(self, rgb):
r, g, b = rgb[:, 0:1, :, :], rgb[:, 1:2, :, :], rgb[:, 2:3, :, :]
gray = 0.2989 * r + 0.5870 * g + 0.1140 * b
return gray
def hamming(self, t1, t2):
dist = (t1 - t2) ** 2
dist_norm = torch.mean(dist / (0.1 + dist), 1, True)
return dist_norm
def valid_mask(self, t, padding):
n, _, h, w = t.size()
inner = torch.ones(n, 1, h - 2 * padding, w - 2 * padding).type_as(t)
mask = F.pad(inner, [padding] * 4)
return mask
def forward(self, img0, img1):
img0 = self.transform(self.rgb2gray(img0))
img1 = self.transform(self.rgb2gray(img1))
return self.hamming(img0, img1) * self.valid_mask(img0, 1)
class SOBEL(nn.Module):
def __init__(self):
super(SOBEL, self).__init__()
self.kernelX = torch.tensor([
[1, 0, -1],
[2, 0, -2],
[1, 0, -1],
]).float()
self.kernelY = self.kernelX.clone().T
self.kernelX = self.kernelX.unsqueeze(0).unsqueeze(0).to(device)
self.kernelY = self.kernelY.unsqueeze(0).unsqueeze(0).to(device)
def forward(self, pred, gt):
N, C, H, W = pred.shape[0], pred.shape[1], pred.shape[2], pred.shape[3]
img_stack = torch.cat(
[pred.reshape(N*C, 1, H, W), gt.reshape(N*C, 1, H, W)], 0)
sobel_stack_x = F.conv2d(img_stack, self.kernelX, padding=1)
sobel_stack_y = F.conv2d(img_stack, self.kernelY, padding=1)
pred_X, gt_X = sobel_stack_x[:N*C], sobel_stack_x[N*C:]
pred_Y, gt_Y = sobel_stack_y[:N*C], sobel_stack_y[N*C:]
L1X, L1Y = torch.abs(pred_X-gt_X), torch.abs(pred_Y-gt_Y)
loss = (L1X+L1Y)
return loss
class MeanShift(nn.Conv2d):
def __init__(self, data_mean, data_std, data_range=1, norm=True):
c = len(data_mean)
super(MeanShift, self).__init__(c, c, kernel_size=1)
std = torch.Tensor(data_std)
self.weight.data = torch.eye(c).view(c, c, 1, 1)
if norm:
self.weight.data.div_(std.view(c, 1, 1, 1))
self.bias.data = -1 * data_range * torch.Tensor(data_mean)
self.bias.data.div_(std)
else:
self.weight.data.mul_(std.view(c, 1, 1, 1))
self.bias.data = data_range * torch.Tensor(data_mean)
self.requires_grad = False
class VGGPerceptualLoss(torch.nn.Module):
def __init__(self, rank=0):
super(VGGPerceptualLoss, self).__init__()
blocks = []
pretrained = True
self.vgg_pretrained_features = models.vgg19(pretrained=pretrained).features
self.normalize = MeanShift([0.485, 0.456, 0.406], [0.229, 0.224, 0.225], norm=True).cuda()
for param in self.parameters():
param.requires_grad = False
def forward(self, X, Y, indices=None):
X = self.normalize(X)
Y = self.normalize(Y)
indices = [2, 7, 12, 21, 30]
weights = [1.0/2.6, 1.0/4.8, 1.0/3.7, 1.0/5.6, 10/1.5]
k = 0
loss = 0
for i in range(indices[-1]):
X = self.vgg_pretrained_features[i](X)
Y = self.vgg_pretrained_features[i](Y)
if (i+1) in indices:
loss += weights[k] * (X - Y.detach()).abs().mean() * 0.1
k += 1
return loss
if __name__ == '__main__':
img0 = torch.zeros(3, 3, 256, 256).float().to(device)
img1 = torch.tensor(np.random.normal(
0, 1, (3, 3, 256, 256))).float().to(device)
ternary_loss = Ternary()
print(ternary_loss(img0, img1).shape)
+200
View File
@@ -0,0 +1,200 @@
import torch
import torch.nn.functional as F
from math import exp
import numpy as np
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def gaussian(window_size, sigma):
gauss = torch.Tensor([exp(-(x - window_size//2)**2/float(2*sigma**2)) for x in range(window_size)])
return gauss/gauss.sum()
def create_window(window_size, channel=1):
_1D_window = gaussian(window_size, 1.5).unsqueeze(1)
_2D_window = _1D_window.mm(_1D_window.t()).float().unsqueeze(0).unsqueeze(0).to(device)
window = _2D_window.expand(channel, 1, window_size, window_size).contiguous()
return window
def create_window_3d(window_size, channel=1):
_1D_window = gaussian(window_size, 1.5).unsqueeze(1)
_2D_window = _1D_window.mm(_1D_window.t())
_3D_window = _2D_window.unsqueeze(2) @ (_1D_window.t())
window = _3D_window.expand(1, channel, window_size, window_size, window_size).contiguous().to(device)
return window
def ssim(img1, img2, window_size=11, window=None, size_average=True, full=False, val_range=None):
# Value range can be different from 255. Other common ranges are 1 (sigmoid) and 2 (tanh).
if val_range is None:
if torch.max(img1) > 128:
max_val = 255
else:
max_val = 1
if torch.min(img1) < -0.5:
min_val = -1
else:
min_val = 0
L = max_val - min_val
else:
L = val_range
padd = 0
(_, channel, height, width) = img1.size()
if window is None:
real_size = min(window_size, height, width)
window = create_window(real_size, channel=channel).to(img1.device)
# mu1 = F.conv2d(img1, window, padding=padd, groups=channel)
# mu2 = F.conv2d(img2, window, padding=padd, groups=channel)
mu1 = F.conv2d(F.pad(img1, (5, 5, 5, 5), mode='replicate'), window, padding=padd, groups=channel)
mu2 = F.conv2d(F.pad(img2, (5, 5, 5, 5), mode='replicate'), window, padding=padd, groups=channel)
mu1_sq = mu1.pow(2)
mu2_sq = mu2.pow(2)
mu1_mu2 = mu1 * mu2
sigma1_sq = F.conv2d(F.pad(img1 * img1, (5, 5, 5, 5), 'replicate'), window, padding=padd, groups=channel) - mu1_sq
sigma2_sq = F.conv2d(F.pad(img2 * img2, (5, 5, 5, 5), 'replicate'), window, padding=padd, groups=channel) - mu2_sq
sigma12 = F.conv2d(F.pad(img1 * img2, (5, 5, 5, 5), 'replicate'), window, padding=padd, groups=channel) - mu1_mu2
C1 = (0.01 * L) ** 2
C2 = (0.03 * L) ** 2
v1 = 2.0 * sigma12 + C2
v2 = sigma1_sq + sigma2_sq + C2
cs = torch.mean(v1 / v2) # contrast sensitivity
ssim_map = ((2 * mu1_mu2 + C1) * v1) / ((mu1_sq + mu2_sq + C1) * v2)
if size_average:
ret = ssim_map.mean()
else:
ret = ssim_map.mean(1).mean(1).mean(1)
if full:
return ret, cs
return ret
def ssim_matlab(img1, img2, window_size=11, window=None, size_average=True, full=False, val_range=None):
# Value range can be different from 255. Other common ranges are 1 (sigmoid) and 2 (tanh).
if val_range is None:
if torch.max(img1) > 128:
max_val = 255
else:
max_val = 1
if torch.min(img1) < -0.5:
min_val = -1
else:
min_val = 0
L = max_val - min_val
else:
L = val_range
padd = 0
(_, _, height, width) = img1.size()
if window is None:
real_size = min(window_size, height, width)
window = create_window_3d(real_size, channel=1).to(img1.device)
# Channel is set to 1 since we consider color images as volumetric images
img1 = img1.unsqueeze(1)
img2 = img2.unsqueeze(1)
mu1 = F.conv3d(F.pad(img1, (5, 5, 5, 5, 5, 5), mode='replicate'), window, padding=padd, groups=1)
mu2 = F.conv3d(F.pad(img2, (5, 5, 5, 5, 5, 5), mode='replicate'), window, padding=padd, groups=1)
mu1_sq = mu1.pow(2)
mu2_sq = mu2.pow(2)
mu1_mu2 = mu1 * mu2
sigma1_sq = F.conv3d(F.pad(img1 * img1, (5, 5, 5, 5, 5, 5), 'replicate'), window, padding=padd, groups=1) - mu1_sq
sigma2_sq = F.conv3d(F.pad(img2 * img2, (5, 5, 5, 5, 5, 5), 'replicate'), window, padding=padd, groups=1) - mu2_sq
sigma12 = F.conv3d(F.pad(img1 * img2, (5, 5, 5, 5, 5, 5), 'replicate'), window, padding=padd, groups=1) - mu1_mu2
C1 = (0.01 * L) ** 2
C2 = (0.03 * L) ** 2
v1 = 2.0 * sigma12 + C2
v2 = sigma1_sq + sigma2_sq + C2
cs = torch.mean(v1 / v2) # contrast sensitivity
ssim_map = ((2 * mu1_mu2 + C1) * v1) / ((mu1_sq + mu2_sq + C1) * v2)
if size_average:
ret = ssim_map.mean()
else:
ret = ssim_map.mean(1).mean(1).mean(1)
if full:
return ret, cs
return ret
def msssim(img1, img2, window_size=11, size_average=True, val_range=None, normalize=False):
device = img1.device
weights = torch.FloatTensor([0.0448, 0.2856, 0.3001, 0.2363, 0.1333]).to(device)
levels = weights.size()[0]
mssim = []
mcs = []
for _ in range(levels):
sim, cs = ssim(img1, img2, window_size=window_size, size_average=size_average, full=True, val_range=val_range)
mssim.append(sim)
mcs.append(cs)
img1 = F.avg_pool2d(img1, (2, 2))
img2 = F.avg_pool2d(img2, (2, 2))
mssim = torch.stack(mssim)
mcs = torch.stack(mcs)
# Normalize (to avoid NaNs during training unstable models, not compliant with original definition)
if normalize:
mssim = (mssim + 1) / 2
mcs = (mcs + 1) / 2
pow1 = mcs ** weights
pow2 = mssim ** weights
# From Matlab implementation https://ece.uwaterloo.ca/~z70wang/research/iwssim/
output = torch.prod(pow1[:-1] * pow2[-1])
return output
# Classes to re-use window
class SSIM(torch.nn.Module):
def __init__(self, window_size=11, size_average=True, val_range=None):
super(SSIM, self).__init__()
self.window_size = window_size
self.size_average = size_average
self.val_range = val_range
# Assume 3 channel for SSIM
self.channel = 3
self.window = create_window(window_size, channel=self.channel)
def forward(self, img1, img2):
(_, channel, _, _) = img1.size()
if channel == self.channel and self.window.dtype == img1.dtype:
window = self.window
else:
window = create_window(self.window_size, channel).to(img1.device).type(img1.dtype)
self.window = window
self.channel = channel
_ssim = ssim(img1, img2, window=window, window_size=self.window_size, size_average=self.size_average)
dssim = (1 - _ssim) / 2
return dssim
class MSSSIM(torch.nn.Module):
def __init__(self, window_size=11, size_average=True, channel=3):
super(MSSSIM, self).__init__()
self.window_size = window_size
self.size_average = size_average
self.channel = channel
def forward(self, img1, img2):
return msssim(img1, img2, window_size=self.window_size, size_average=self.size_average)
+82
View File
@@ -0,0 +1,82 @@
import torch
import torch.nn as nn
import numpy as np
import torch.optim as optim
import itertools
from model.warplayer import warp
import torch.nn.functional as F
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation, bias=True),
nn.PReLU(out_planes)
)
def deconv(in_planes, out_planes, kernel_size=4, stride=2, padding=1):
return nn.Sequential(
torch.nn.ConvTranspose2d(in_channels=in_planes, out_channels=out_planes, kernel_size=4, stride=2, padding=1, bias=True),
nn.PReLU(out_planes)
)
class Conv2(nn.Module):
def __init__(self, in_planes, out_planes, stride=2):
super(Conv2, self).__init__()
self.conv1 = conv(in_planes, out_planes, 3, stride, 1)
self.conv2 = conv(out_planes, out_planes, 3, 1, 1)
def forward(self, x):
x = self.conv1(x)
x = self.conv2(x)
return x
c = 16
class Contextnet(nn.Module):
def __init__(self):
super(Contextnet, self).__init__()
self.conv1 = Conv2(3, c)
self.conv2 = Conv2(c, 2*c)
self.conv3 = Conv2(2*c, 4*c)
self.conv4 = Conv2(4*c, 8*c)
def forward(self, x, flow):
x = self.conv1(x)
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 0.5
f1 = warp(x, flow)
x = self.conv2(x)
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 0.5
f2 = warp(x, flow)
x = self.conv3(x)
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 0.5
f3 = warp(x, flow)
x = self.conv4(x)
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 0.5
f4 = warp(x, flow)
return [f1, f2, f3, f4]
class Unet(nn.Module):
def __init__(self):
super(Unet, self).__init__()
self.down0 = Conv2(17, 2*c)
self.down1 = Conv2(4*c, 4*c)
self.down2 = Conv2(8*c, 8*c)
self.down3 = Conv2(16*c, 16*c)
self.up0 = deconv(32*c, 8*c)
self.up1 = deconv(16*c, 4*c)
self.up2 = deconv(8*c, 2*c)
self.up3 = deconv(4*c, c)
self.conv = nn.Conv2d(c, 3, 3, 1, 1)
def forward(self, img0, img1, warped_img0, warped_img1, mask, flow, c0, c1):
s0 = self.down0(torch.cat((img0, img1, warped_img0, warped_img1, mask, flow), 1))
s1 = self.down1(torch.cat((s0, c0[0], c1[0]), 1))
s2 = self.down2(torch.cat((s1, c0[1], c1[1]), 1))
s3 = self.down3(torch.cat((s2, c0[2], c1[2]), 1))
x = self.up0(torch.cat((s3, c0[3], c1[3]), 1))
x = self.up1(torch.cat((x, s2), 1))
x = self.up2(torch.cat((x, s1), 1))
x = self.up3(torch.cat((x, s0), 1))
x = self.conv(x)
return torch.sigmoid(x)
+83
View File
@@ -0,0 +1,83 @@
import torch
import torch.nn as nn
import numpy as np
import torch.optim as optim
import itertools
from model.warplayer import warp
from torch.nn.parallel import DistributedDataParallel as DDP
import torch.nn.functional as F
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def conv(in_planes, out_planes, kernel_size=3, stride=1, padding=1, dilation=1):
return nn.Sequential(
nn.Conv2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride,
padding=padding, dilation=dilation, bias=True),
nn.PReLU(out_planes)
)
def deconv(in_planes, out_planes, kernel_size=4, stride=2, padding=1):
return nn.Sequential(
torch.nn.ConvTranspose2d(in_channels=in_planes, out_channels=out_planes, kernel_size=4, stride=2, padding=1, bias=True),
nn.PReLU(out_planes)
)
class Conv2(nn.Module):
def __init__(self, in_planes, out_planes, stride=2):
super(Conv2, self).__init__()
self.conv1 = conv(in_planes, out_planes, 3, stride, 1)
self.conv2 = conv(out_planes, out_planes, 3, 1, 1)
def forward(self, x):
x = self.conv1(x)
x = self.conv2(x)
return x
c = 16
class Contextnet(nn.Module):
def __init__(self):
super(Contextnet, self).__init__()
self.conv1 = Conv2(3, c, 1)
self.conv2 = Conv2(c, 2*c)
self.conv3 = Conv2(2*c, 4*c)
self.conv4 = Conv2(4*c, 8*c)
def forward(self, x, flow):
x = self.conv1(x)
# flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 0.5
f1 = warp(x, flow)
x = self.conv2(x)
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 0.5
f2 = warp(x, flow)
x = self.conv3(x)
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 0.5
f3 = warp(x, flow)
x = self.conv4(x)
flow = F.interpolate(flow, scale_factor=0.5, mode="bilinear", align_corners=False, recompute_scale_factor=False) * 0.5
f4 = warp(x, flow)
return [f1, f2, f3, f4]
class Unet(nn.Module):
def __init__(self):
super(Unet, self).__init__()
self.down0 = Conv2(17, 2*c, 1)
self.down1 = Conv2(4*c, 4*c)
self.down2 = Conv2(8*c, 8*c)
self.down3 = Conv2(16*c, 16*c)
self.up0 = deconv(32*c, 8*c)
self.up1 = deconv(16*c, 4*c)
self.up2 = deconv(8*c, 2*c)
self.up3 = deconv(4*c, c)
self.conv = nn.Conv2d(c, 3, 3, 2, 1)
def forward(self, img0, img1, warped_img0, warped_img1, mask, flow, c0, c1):
s0 = self.down0(torch.cat((img0, img1, warped_img0, warped_img1, mask, flow), 1))
s1 = self.down1(torch.cat((s0, c0[0], c1[0]), 1))
s2 = self.down2(torch.cat((s1, c0[1], c1[1]), 1))
s3 = self.down3(torch.cat((s2, c0[2], c1[2]), 1))
x = self.up0(torch.cat((s3, c0[3], c1[3]), 1))
x = self.up1(torch.cat((x, s2), 1))
x = self.up2(torch.cat((x, s1), 1))
x = self.up3(torch.cat((x, s0), 1))
x = self.conv(x)
return torch.sigmoid(x)
+22
View File
@@ -0,0 +1,22 @@
import torch
import torch.nn as nn
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
backwarp_tenGrid = {}
def warp(tenInput, tenFlow):
k = (str(tenFlow.device), str(tenFlow.size()))
if k not in backwarp_tenGrid:
tenHorizontal = torch.linspace(-1.0, 1.0, tenFlow.shape[3], device=device).view(
1, 1, 1, tenFlow.shape[3]).expand(tenFlow.shape[0], -1, tenFlow.shape[2], -1)
tenVertical = torch.linspace(-1.0, 1.0, tenFlow.shape[2], device=device).view(
1, 1, tenFlow.shape[2], 1).expand(tenFlow.shape[0], -1, -1, tenFlow.shape[3])
backwarp_tenGrid[k] = torch.cat(
[tenHorizontal, tenVertical], 1).to(device)
tenFlow = torch.cat([tenFlow[:, 0:1, :, :] / ((tenInput.shape[3] - 1.0) / 2.0),
tenFlow[:, 1:2, :, :] / ((tenInput.shape[2] - 1.0) / 2.0)], 1)
g = (backwarp_tenGrid[k] + tenFlow).permute(0, 2, 3, 1)
return torch.nn.functional.grid_sample(input=tenInput, grid=g, mode='bilinear', padding_mode='border', align_corners=True)
BIN
View File
Binary file not shown.