Init
This commit is contained in:
@@ -0,0 +1 @@
|
||||
__pycache__
|
||||
+144
@@ -0,0 +1,144 @@
|
||||
import torch
|
||||
from comfy import model_management
|
||||
import folder_paths
|
||||
from .evtexture.evtexture_arch import EvTexture
|
||||
from .esim import events_generator, events_to_image, EventSimulatorConfig
|
||||
from .evoxels import package_bidirectional_event_voxels
|
||||
|
||||
EVENTS_TYPE = "EVT_EVENTS"
|
||||
EVTEXTURE_MODEL_TYPE = "EVTEXTURE_MODEL"
|
||||
|
||||
|
||||
class VideoToEvents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {"images": ("IMAGE", {}), "fps": ("FLOAT", {})},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (EVENTS_TYPE,)
|
||||
RETURN_NAMES = ("events",)
|
||||
CATEGORY = "EVTexture"
|
||||
FUNCTION = "events"
|
||||
|
||||
def events(self, images, fps: float):
|
||||
imgs = torch.mean(images, dim=3)
|
||||
log_imgs = (imgs + 1e-3).log()
|
||||
|
||||
timestamps = [i / fps for i in range(len(images))]
|
||||
config = EventSimulatorConfig()
|
||||
return (list(events_generator(log_imgs, timestamps, config)),)
|
||||
|
||||
|
||||
class EventsToImage:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"events": (EVENTS_TYPE, {"forceInput": True}),
|
||||
"width": ("INT", {}),
|
||||
"height": ("INT", {}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("images",)
|
||||
CATEGORY = "EVTexture"
|
||||
FUNCTION = "to_image"
|
||||
|
||||
def to_image(self, events, height: int, width: int):
|
||||
b, h, w = len(events), height, width
|
||||
res = torch.zeros((b, h, w, 3))
|
||||
for i in range(b):
|
||||
res[i, :, :, :] = events_to_image(events[i], h, w).permute(1, 2, 0)
|
||||
return (res,)
|
||||
|
||||
|
||||
class LoadEvTextureModel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model_name": (folder_paths.get_filename_list("upscale_models"),),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (EVTEXTURE_MODEL_TYPE,)
|
||||
RETURN_NAMES = ("model",)
|
||||
CATEGORY = "EVTexture"
|
||||
FUNCTION = "load"
|
||||
|
||||
def load(self, model_name):
|
||||
path = folder_paths.get_full_path_or_raise("upscale_models", model_name)
|
||||
model = EvTexture()
|
||||
params_dict = torch.load(path, weights_only=True)["params_ema"]
|
||||
model.load_state_dict(params_dict, strict=True)
|
||||
return (model,)
|
||||
|
||||
|
||||
class EvTextureUpscaleVideo:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE", {}),
|
||||
"events": (EVENTS_TYPE, {}),
|
||||
"model": (EVTEXTURE_MODEL_TYPE, {}),
|
||||
"fps": ("FLOAT", {}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("images",)
|
||||
CATEGORY = "EVTexture"
|
||||
FUNCTION = "upscale"
|
||||
|
||||
def upscale(self, images, events, model: EvTexture, fps: float):
|
||||
device = model_management.get_torch_device()
|
||||
|
||||
n, h, w, _ = images.shape
|
||||
memory_required = model_management.module_size(model)
|
||||
memory_required += (
|
||||
n * (h * w * 3) * images.element_size() * 4 * 384.0
|
||||
) # The 384.0 is an estimate of how much some of these models take, TODO: make it more accurate
|
||||
memory_required += images.nelement() * images.element_size()
|
||||
model_management.free_memory(memory_required, device)
|
||||
|
||||
model.to(device)
|
||||
imgs = images.movedim(-1, -3).unsqueeze(0).to(device)
|
||||
events = torch.vstack(events).to(device)
|
||||
|
||||
xs, ys, ts, pols = events.T
|
||||
timestamps = [i / fps for i in range(n)]
|
||||
bins = 5
|
||||
voxels_f = torch.stack(
|
||||
package_bidirectional_event_voxels(
|
||||
xs, ys, ts, pols, timestamps, False, bins, (h, w)
|
||||
)
|
||||
).unsqueeze(0)
|
||||
voxels_b = torch.stack(
|
||||
package_bidirectional_event_voxels(
|
||||
xs, ys, ts, pols, timestamps, True, bins, (h, w)
|
||||
)
|
||||
).unsqueeze(0)
|
||||
del xs, ys, ts, pols, events
|
||||
|
||||
s = model.forward(imgs, voxels_f, voxels_b)[0].to("cpu")
|
||||
|
||||
model.to("cpu")
|
||||
s = torch.clamp(s.movedim(-3, -1), min=0, max=1.0)
|
||||
return (s,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"EVTVideoToEvents": VideoToEvents,
|
||||
"EVTEventsToImage": EventsToImage,
|
||||
"EVTLoadEvTextureModel": LoadEvTextureModel,
|
||||
"EVTTextureUpscaleVideo": EvTextureUpscaleVideo,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"EVTVideoToEvents": "Video to Camera Events",
|
||||
"EVTEventsToImage": "Camera Events To Images",
|
||||
"EVTLoadEvTextureModel": "Load EvTexture Model",
|
||||
"EVTTextureUpscaleVideo": "EvTexture Video Upscale",
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
import torch
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass
|
||||
class EventSimulatorConfig:
|
||||
contrast_threshold_pos: float = 0.275
|
||||
contrast_threshold_neg: float = 0.275
|
||||
refractory_period: float = 1e-4
|
||||
|
||||
|
||||
def events(
|
||||
img: torch.Tensor,
|
||||
time: float,
|
||||
last_img: torch.Tensor,
|
||||
last_time: float,
|
||||
last_event_timestamp: torch.Tensor,
|
||||
ref_values: torch.Tensor,
|
||||
config: EventSimulatorConfig,
|
||||
):
|
||||
delta = img - last_img
|
||||
delta_t = time - last_time
|
||||
|
||||
pol = torch.where(delta >= 0.0, 1.0, -1.0)
|
||||
contrast_threshold = torch.where(
|
||||
pol > 0, config.contrast_threshold_pos, config.contrast_threshold_neg
|
||||
)
|
||||
|
||||
active_mask = delta.abs() > 1e-6
|
||||
events = []
|
||||
curr_cross = ref_values.clone()
|
||||
while True:
|
||||
curr_cross[active_mask] += pol[active_mask] * contrast_threshold[active_mask]
|
||||
pos_crossing = (pol > 0) & (curr_cross > last_img) & (curr_cross <= img)
|
||||
neg_crossing = (pol < 0) & (curr_cross < last_img) & (curr_cross >= img)
|
||||
crossing_conditions = pos_crossing | neg_crossing
|
||||
|
||||
active_mask &= crossing_conditions
|
||||
|
||||
if not active_mask.any(): # loop until no activations
|
||||
break
|
||||
|
||||
ref_values[active_mask] = curr_cross[active_mask]
|
||||
|
||||
edt = torch.zeros_like(img)
|
||||
edt[active_mask] = (
|
||||
(curr_cross[active_mask] - last_img[active_mask])
|
||||
* delta_t
|
||||
/ delta[active_mask]
|
||||
)
|
||||
t = last_time + edt
|
||||
|
||||
event_mask = (t - last_event_timestamp) >= config.refractory_period
|
||||
event_mask |= last_event_timestamp == 0
|
||||
event_mask &= active_mask
|
||||
last_event_timestamp[event_mask] = t[event_mask]
|
||||
|
||||
indices = event_mask.argwhere()
|
||||
ys, xs = indices[:, 0], indices[:, 1]
|
||||
events.append(torch.column_stack((xs, ys, t[ys, xs], pol[ys, xs])))
|
||||
|
||||
if events:
|
||||
events = torch.vstack(events)
|
||||
else:
|
||||
events = torch.empty((0, 4))
|
||||
events = events[events[:, 2].argsort()]
|
||||
return events
|
||||
|
||||
|
||||
def events_generator(imgs, timestamps, config: EventSimulatorConfig):
|
||||
it = iter(zip(imgs, timestamps))
|
||||
last_img, last_time = next(it)
|
||||
last_img = last_img.squeeze()
|
||||
|
||||
assert len(last_img.shape) == 2, "expected single channel images of shape [h, w]"
|
||||
|
||||
last_event_timestamp = torch.zeros_like(last_img)
|
||||
ref_values = last_img.clone()
|
||||
|
||||
for img, time in it:
|
||||
img = img.squeeze()
|
||||
yield events(
|
||||
img, time, last_img, last_time, last_event_timestamp, ref_values, config
|
||||
)
|
||||
last_img = img
|
||||
last_time = time
|
||||
|
||||
|
||||
def events_to_image(events, height: int, width: int):
|
||||
img = torch.zeros((3, height, width), dtype=torch.float32)
|
||||
xs = events[:, 0].long()
|
||||
ys = events[:, 1].long()
|
||||
pol = events[:, 3]
|
||||
pos_mask = pol > 0
|
||||
neg_mask = pol < 0
|
||||
img[0, ys[pos_mask], xs[pos_mask]] = 1.0
|
||||
img[2, ys[neg_mask], xs[neg_mask]] = 1.0
|
||||
return img
|
||||
+262
@@ -0,0 +1,262 @@
|
||||
import torch
|
||||
import math
|
||||
import bisect
|
||||
|
||||
|
||||
## Extracted from https://github.com/TimoStoff/event_utils
|
||||
def events_to_voxel_torch(
|
||||
xs, ys, ts, ps, B, device=None, sensor_size=(180, 240), temporal_bilinear=True
|
||||
):
|
||||
"""
|
||||
Turn set of events to a voxel grid tensor, using temporal bilinear interpolation
|
||||
@param xs List of event x coordinates (torch tensor)
|
||||
@param ys List of event y coordinates (torch tensor)
|
||||
@param ts List of event timestamps (torch tensor)
|
||||
@param ps List of event polarities (torch tensor)
|
||||
@param B Number of bins in output voxel grids (int)
|
||||
@param device Device to put voxel grid. If left empty, same device as events
|
||||
@param sensor_size The size of the event sensor/output voxels
|
||||
@param temporal_bilinear Whether the events should be naively
|
||||
accumulated to the voxels (faster), or properly
|
||||
temporally distributed
|
||||
@returns Voxel of the events between t0 and t1
|
||||
"""
|
||||
if device is None:
|
||||
device = xs.device
|
||||
assert len(xs) == len(ys) and len(ys) == len(ts) and len(ts) == len(ps)
|
||||
bins = []
|
||||
dt = ts[-1] - ts[0]
|
||||
t_norm = (ts - ts[0]) / dt * (B - 1)
|
||||
zeros = torch.zeros_like(t_norm)
|
||||
for bi in range(B):
|
||||
assert temporal_bilinear, "no other option not supported"
|
||||
bilinear_weights = torch.max(zeros, 1.0 - torch.abs(t_norm - bi))
|
||||
weights = ps * bilinear_weights
|
||||
vb = events_to_image_torch(
|
||||
xs,
|
||||
ys,
|
||||
weights,
|
||||
device,
|
||||
sensor_size=sensor_size,
|
||||
clip_out_of_range=False,
|
||||
)
|
||||
bins.append(vb)
|
||||
bins = torch.stack(bins)
|
||||
return bins
|
||||
|
||||
|
||||
## Extracted from https://github.com/TimoStoff/event_utils
|
||||
def events_to_image_torch(
|
||||
xs,
|
||||
ys,
|
||||
ps,
|
||||
device=None,
|
||||
sensor_size=(180, 240),
|
||||
clip_out_of_range=True,
|
||||
interpolation=None,
|
||||
padding=True,
|
||||
default=0,
|
||||
):
|
||||
"""
|
||||
Method to turn event tensor to image. Allows for bilinear interpolation.
|
||||
@param xs Tensor of x coords of events
|
||||
@param ys Tensor of y coords of events
|
||||
@param ps Tensor of event polarities/weights
|
||||
@param device The device on which the image is. If none, set to events device
|
||||
@param sensor_size The size of the image sensor/output image
|
||||
@param clip_out_of_range If the events go beyond the desired image size,
|
||||
clip the events to fit into the image
|
||||
@param interpolation Which interpolation to use. Options=None,'bilinear'
|
||||
@param padding If bilinear interpolation, allow padding the image by 1 to allow events to fit:
|
||||
@returns Event image from the events
|
||||
"""
|
||||
if device is None:
|
||||
device = xs.device
|
||||
if interpolation == "bilinear" and padding:
|
||||
img_size = (sensor_size[0] + 1, sensor_size[1] + 1)
|
||||
else:
|
||||
img_size = list(sensor_size)
|
||||
|
||||
mask = torch.ones(xs.size(), device=device)
|
||||
if clip_out_of_range:
|
||||
zero_v = torch.tensor([0.0], device=device)
|
||||
ones_v = torch.tensor([1.0], device=device)
|
||||
clipx = (
|
||||
img_size[1]
|
||||
if interpolation is None and padding == False
|
||||
else img_size[1] - 1
|
||||
)
|
||||
clipy = (
|
||||
img_size[0]
|
||||
if interpolation is None and padding == False
|
||||
else img_size[0] - 1
|
||||
)
|
||||
mask = torch.where(xs >= clipx, zero_v, ones_v) * torch.where(
|
||||
ys >= clipy, zero_v, ones_v
|
||||
)
|
||||
|
||||
img = (torch.ones(img_size) * default).to(device)
|
||||
if (
|
||||
interpolation == "bilinear"
|
||||
and xs.dtype is not torch.long
|
||||
and xs.dtype is not torch.long
|
||||
):
|
||||
pxs = (xs.floor()).float()
|
||||
pys = (ys.floor()).float()
|
||||
dxs = (xs - pxs).float()
|
||||
dys = (ys - pys).float()
|
||||
pxs = (pxs * mask).long()
|
||||
pys = (pys * mask).long()
|
||||
masked_ps = ps.squeeze() * mask
|
||||
interpolate_to_image(pxs, pys, dxs, dys, masked_ps, img)
|
||||
else:
|
||||
if xs.dtype is not torch.long:
|
||||
xs = xs.long().to(device)
|
||||
if ys.dtype is not torch.long:
|
||||
ys = ys.long().to(device)
|
||||
try:
|
||||
mask = mask.long().to(device)
|
||||
xs, ys = xs * mask, ys * mask
|
||||
img.index_put_((ys, xs), ps, accumulate=True)
|
||||
except Exception as e:
|
||||
print(
|
||||
"Unable to put tensor {} positions ({}, {}) into {}. Range = {},{}".format(
|
||||
ps.shape,
|
||||
ys.shape,
|
||||
xs.shape,
|
||||
img.shape,
|
||||
torch.max(ys),
|
||||
torch.max(xs),
|
||||
)
|
||||
)
|
||||
raise e
|
||||
return img
|
||||
|
||||
|
||||
## Extracted from https://github.com/TimoStoff/event_utils
|
||||
def interpolate_to_image(pxs, pys, dxs, dys, weights, img):
|
||||
"""
|
||||
Accumulate x and y coords to an image using bilinear interpolation
|
||||
@param pxs Numpy array of integer typecast x coords of events
|
||||
@param pys Numpy array of integer typecast y coords of events
|
||||
@param dxs Numpy array of residual difference between x coord and int(x coord)
|
||||
@param dys Numpy array of residual difference between y coord and int(y coord)
|
||||
@returns Image
|
||||
"""
|
||||
img.index_put_((pys, pxs), weights * (1.0 - dxs) * (1.0 - dys), accumulate=True)
|
||||
img.index_put_((pys, pxs + 1), weights * dxs * (1.0 - dys), accumulate=True)
|
||||
img.index_put_((pys + 1, pxs), weights * (1.0 - dxs) * dys, accumulate=True)
|
||||
img.index_put_((pys + 1, pxs + 1), weights * dxs * dys, accumulate=True)
|
||||
return img
|
||||
|
||||
|
||||
def voxel_normalization(voxel):
|
||||
"""
|
||||
normalize the voxel same as https://arxiv.org/abs/1912.01584 Section 3.1
|
||||
Params:
|
||||
voxel: torch.Tensor, shape is [num_bins, H, W]
|
||||
|
||||
return:
|
||||
normalized voxel
|
||||
"""
|
||||
# check if voxel all element is 0
|
||||
tmp = torch.zeros_like(voxel)
|
||||
if torch.equal(voxel, tmp):
|
||||
return voxel
|
||||
abs_voxel, _ = torch.sort(torch.abs(voxel).view(-1, 1).squeeze(1))
|
||||
first_non_zero_idx = torch.nonzero(abs_voxel)[0].item()
|
||||
non_zero_voxel = abs_voxel[first_non_zero_idx:]
|
||||
norm_idx = math.floor(non_zero_voxel.shape[0] * 0.98)
|
||||
ones = torch.ones_like(voxel)
|
||||
normed_voxel = torch.where(
|
||||
torch.abs(voxel) < non_zero_voxel[norm_idx],
|
||||
voxel / non_zero_voxel[norm_idx],
|
||||
voxel,
|
||||
)
|
||||
normed_voxel = torch.where(
|
||||
normed_voxel >= non_zero_voxel[norm_idx], ones, normed_voxel
|
||||
)
|
||||
normed_voxel = torch.where(
|
||||
normed_voxel <= -non_zero_voxel[norm_idx], -ones, normed_voxel
|
||||
)
|
||||
return normed_voxel
|
||||
|
||||
|
||||
# Taken and modified from https://github.com/DachunKai/EvTexture/issues/12#issuecomment-2198243470
|
||||
def package_bidirectional_event_voxels(
|
||||
x,
|
||||
y,
|
||||
t,
|
||||
p,
|
||||
timestamp_list,
|
||||
backward,
|
||||
bins,
|
||||
sensor_size,
|
||||
):
|
||||
"""
|
||||
params:
|
||||
x: ndarray, x-position of events
|
||||
y: ndarray, y-position of events
|
||||
t: ndarray, timestamp of events
|
||||
p: ndarray, polarity of events
|
||||
backward: bool, if forward or backward
|
||||
timestamp_list: list, to split events via timestamp
|
||||
bins: voxel num_bins
|
||||
returns:
|
||||
no return.
|
||||
"""
|
||||
assert x.shape == y.shape == t.shape == p.shape
|
||||
|
||||
# Step 2: select events between two frames according to timestamp
|
||||
temp = t.cpu().numpy().tolist()
|
||||
output = [
|
||||
temp[
|
||||
bisect.bisect_left(temp, timestamp_list[i]) : bisect.bisect_left(
|
||||
temp, timestamp_list[i + 1]
|
||||
)
|
||||
]
|
||||
for i in range(len(timestamp_list) - 1)
|
||||
]
|
||||
|
||||
# Debug: Check if data error!!!
|
||||
assert (
|
||||
len(output) == len(timestamp_list) - 1
|
||||
), f"len(output) is {len(output)}, but len(timestamp_list) is {len(timestamp_list)}"
|
||||
sum_output = []
|
||||
sum = 0
|
||||
for i in range(len(output)):
|
||||
if len(output[i]) <= 1:
|
||||
raise ValueError(f"len(output[{i}] == 0)")
|
||||
sum += len(output[i])
|
||||
sum_output.append(sum)
|
||||
|
||||
assert len(sum_output) == len(output)
|
||||
|
||||
# Step 3: After checking data, continue.
|
||||
start_idx = 0
|
||||
out_voxels = []
|
||||
for voxel_idx in range(len(timestamp_list) - 1):
|
||||
end_idx = start_idx + len(output[voxel_idx])
|
||||
|
||||
xs = x[start_idx:end_idx]
|
||||
ys = y[start_idx:end_idx]
|
||||
ts = t[start_idx:end_idx]
|
||||
ps = p[start_idx:end_idx]
|
||||
|
||||
if backward:
|
||||
t_start = timestamp_list[voxel_idx]
|
||||
t_end = timestamp_list[voxel_idx + 1]
|
||||
xs = torch.flip(xs, dims=[0])
|
||||
ys = torch.flip(ys, dims=[0])
|
||||
ts = torch.flip(t_end - ts + t_start, dims=[0])
|
||||
ps = torch.flip(-ps, dims=[0])
|
||||
|
||||
voxel = events_to_voxel_torch(
|
||||
xs, ys, ts, ps, bins, device=None, sensor_size=sensor_size
|
||||
)
|
||||
normed_voxel = voxel_normalization(voxel)
|
||||
|
||||
out_voxels.append(normed_voxel)
|
||||
start_idx = end_idx
|
||||
|
||||
return out_voxels
|
||||
@@ -0,0 +1,439 @@
|
||||
import collections.abc
|
||||
import math
|
||||
import torch
|
||||
import torchvision
|
||||
import warnings
|
||||
from itertools import repeat
|
||||
from torch import nn as nn
|
||||
from torch.nn import functional as F
|
||||
from torch.nn import init as init
|
||||
from torch.nn.modules.batchnorm import _BatchNorm
|
||||
from torch.autograd import Function
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def default_init_weights(module_list, scale=1, bias_fill=0, **kwargs):
|
||||
"""Initialize network weights.
|
||||
|
||||
Args:
|
||||
module_list (list[nn.Module] | nn.Module): Modules to be initialized.
|
||||
scale (float): Scale initialized weights, especially for residual
|
||||
blocks. Default: 1.
|
||||
bias_fill (float): The value to fill bias. Default: 0
|
||||
kwargs (dict): Other arguments for initialization function.
|
||||
"""
|
||||
if not isinstance(module_list, list):
|
||||
module_list = [module_list]
|
||||
for module in module_list:
|
||||
for m in module.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
init.kaiming_normal_(m.weight, **kwargs)
|
||||
m.weight.data *= scale
|
||||
if m.bias is not None:
|
||||
m.bias.data.fill_(bias_fill)
|
||||
elif isinstance(m, nn.Linear):
|
||||
init.kaiming_normal_(m.weight, **kwargs)
|
||||
m.weight.data *= scale
|
||||
if m.bias is not None:
|
||||
m.bias.data.fill_(bias_fill)
|
||||
elif isinstance(m, _BatchNorm):
|
||||
init.constant_(m.weight, 1)
|
||||
if m.bias is not None:
|
||||
m.bias.data.fill_(bias_fill)
|
||||
|
||||
|
||||
def make_layer(basic_block, num_basic_block, **kwarg):
|
||||
"""Make layers by stacking the same blocks.
|
||||
|
||||
Args:
|
||||
basic_block (nn.module): nn.module class for basic block.
|
||||
num_basic_block (int): number of blocks.
|
||||
|
||||
Returns:
|
||||
nn.Sequential: Stacked blocks in nn.Sequential.
|
||||
"""
|
||||
layers = []
|
||||
for _ in range(num_basic_block):
|
||||
layers.append(basic_block(**kwarg))
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
|
||||
class ResidualBlockNoBN(nn.Module):
|
||||
"""Residual block without BN.
|
||||
|
||||
Args:
|
||||
num_feat (int): Channel number of intermediate features.
|
||||
Default: 64.
|
||||
res_scale (float): Residual scale. Default: 1.
|
||||
pytorch_init (bool): If set to True, use pytorch default init,
|
||||
otherwise, use default_init_weights. Default: False.
|
||||
"""
|
||||
|
||||
def __init__(self, num_feat=64, res_scale=1, pytorch_init=False):
|
||||
super(ResidualBlockNoBN, self).__init__()
|
||||
self.res_scale = res_scale
|
||||
self.conv1 = nn.Conv2d(num_feat, num_feat, 3, 1, 1, bias=True)
|
||||
self.conv2 = nn.Conv2d(num_feat, num_feat, 3, 1, 1, bias=True)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
|
||||
if not pytorch_init:
|
||||
default_init_weights([self.conv1, self.conv2], 0.1)
|
||||
|
||||
def forward(self, x):
|
||||
identity = x
|
||||
out = self.conv2(self.relu(self.conv1(x)))
|
||||
return identity + out * self.res_scale
|
||||
|
||||
|
||||
class Upsample(nn.Sequential):
|
||||
"""Upsample module.
|
||||
|
||||
Args:
|
||||
scale (int): Scale factor. Supported scales: 2^n and 3.
|
||||
num_feat (int): Channel number of intermediate features.
|
||||
"""
|
||||
|
||||
def __init__(self, scale, num_feat):
|
||||
m = []
|
||||
if (scale & (scale - 1)) == 0: # scale = 2^n
|
||||
for _ in range(int(math.log(scale, 2))):
|
||||
m.append(nn.Conv2d(num_feat, 4 * num_feat, 3, 1, 1))
|
||||
m.append(nn.PixelShuffle(2))
|
||||
elif scale == 3:
|
||||
m.append(nn.Conv2d(num_feat, 9 * num_feat, 3, 1, 1))
|
||||
m.append(nn.PixelShuffle(3))
|
||||
else:
|
||||
raise ValueError(
|
||||
f"scale {scale} is not supported. Supported scales: 2^n and 3."
|
||||
)
|
||||
super(Upsample, self).__init__(*m)
|
||||
|
||||
|
||||
def flow_warp(
|
||||
x, flow, interp_mode="bilinear", padding_mode="zeros", align_corners=True
|
||||
):
|
||||
"""Warp an image or feature map with optical flow.
|
||||
|
||||
Args:
|
||||
x (Tensor): Tensor with size (n, c, h, w).
|
||||
flow (Tensor): Tensor with size (n, h, w, 2), normal value.
|
||||
interp_mode (str): 'nearest' or 'bilinear'. Default: 'bilinear'.
|
||||
padding_mode (str): 'zeros' or 'border' or 'reflection'.
|
||||
Default: 'zeros'.
|
||||
align_corners (bool): Before pytorch 1.3, the default value is
|
||||
align_corners=True. After pytorch 1.3, the default value is
|
||||
align_corners=False. Here, we use the True as default.
|
||||
|
||||
Returns:
|
||||
Tensor: Warped image or feature map.
|
||||
"""
|
||||
assert x.size()[-2:] == flow.size()[1:3]
|
||||
_, _, h, w = x.size()
|
||||
# create mesh grid
|
||||
grid_y, grid_x = torch.meshgrid(
|
||||
torch.arange(0, h, dtype=x.dtype, device=x.device),
|
||||
torch.arange(0, w, dtype=x.dtype, device=x.device),
|
||||
indexing="ij",
|
||||
)
|
||||
grid = torch.stack((grid_x, grid_y), 2).float() # W(x), H(y), 2
|
||||
grid.requires_grad = False
|
||||
|
||||
vgrid = grid + flow
|
||||
# scale grid to [-1,1]
|
||||
vgrid_x = 2.0 * vgrid[:, :, :, 0] / max(w - 1, 1) - 1.0
|
||||
vgrid_y = 2.0 * vgrid[:, :, :, 1] / max(h - 1, 1) - 1.0
|
||||
vgrid_scaled = torch.stack((vgrid_x, vgrid_y), dim=3)
|
||||
output = F.grid_sample(
|
||||
x,
|
||||
vgrid_scaled,
|
||||
mode=interp_mode,
|
||||
padding_mode=padding_mode,
|
||||
align_corners=align_corners,
|
||||
)
|
||||
|
||||
# TODO, what if align_corners=False
|
||||
return output
|
||||
|
||||
|
||||
def resize_flow(flow, size_type, sizes, interp_mode="bilinear", align_corners=False):
|
||||
"""Resize a flow according to ratio or shape.
|
||||
|
||||
Args:
|
||||
flow (Tensor): Precomputed flow. shape [N, 2, H, W].
|
||||
size_type (str): 'ratio' or 'shape'.
|
||||
sizes (list[int | float]): the ratio for resizing or the final output
|
||||
shape.
|
||||
1) The order of ratio should be [ratio_h, ratio_w]. For
|
||||
downsampling, the ratio should be smaller than 1.0 (i.e., ratio
|
||||
< 1.0). For upsampling, the ratio should be larger than 1.0 (i.e.,
|
||||
ratio > 1.0).
|
||||
2) The order of output_size should be [out_h, out_w].
|
||||
interp_mode (str): The mode of interpolation for resizing.
|
||||
Default: 'bilinear'.
|
||||
align_corners (bool): Whether align corners. Default: False.
|
||||
|
||||
Returns:
|
||||
Tensor: Resized flow.
|
||||
"""
|
||||
_, _, flow_h, flow_w = flow.size()
|
||||
if size_type == "ratio":
|
||||
output_h, output_w = int(flow_h * sizes[0]), int(flow_w * sizes[1])
|
||||
elif size_type == "shape":
|
||||
output_h, output_w = sizes[0], sizes[1]
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Size type should be ratio or shape, but got type {size_type}."
|
||||
)
|
||||
|
||||
input_flow = flow.clone()
|
||||
ratio_h = output_h / flow_h
|
||||
ratio_w = output_w / flow_w
|
||||
input_flow[:, 0, :, :] *= ratio_w
|
||||
input_flow[:, 1, :, :] *= ratio_h
|
||||
resized_flow = F.interpolate(
|
||||
input=input_flow,
|
||||
size=(output_h, output_w),
|
||||
mode=interp_mode,
|
||||
align_corners=align_corners,
|
||||
)
|
||||
return resized_flow
|
||||
|
||||
|
||||
# TODO: may write a cpp file
|
||||
def pixel_unshuffle(x, scale):
|
||||
"""Pixel unshuffle.
|
||||
|
||||
Args:
|
||||
x (Tensor): Input feature with shape (b, c, hh, hw).
|
||||
scale (int): Downsample ratio.
|
||||
|
||||
Returns:
|
||||
Tensor: the pixel unshuffled feature.
|
||||
"""
|
||||
b, c, hh, hw = x.size()
|
||||
out_channel = c * (scale**2)
|
||||
assert hh % scale == 0 and hw % scale == 0
|
||||
h = hh // scale
|
||||
w = hw // scale
|
||||
x_view = x.view(b, c, h, scale, w, scale)
|
||||
return x_view.permute(0, 1, 3, 5, 2, 4).reshape(b, out_channel, h, w)
|
||||
|
||||
|
||||
class DCNv2Pack(Function):
|
||||
"""Modulated deformable conv for deformable alignment.
|
||||
|
||||
Different from the official DCNv2Pack, which generates offsets and masks
|
||||
from the preceding features, this DCNv2Pack takes another different
|
||||
features to generate offsets and masks.
|
||||
|
||||
``Paper: Delving Deep into Deformable Alignment in Video Super-Resolution``
|
||||
"""
|
||||
|
||||
def forward(self, x, feat):
|
||||
out = self.conv_offset(feat)
|
||||
o1, o2, mask = torch.chunk(out, 3, dim=1)
|
||||
offset = torch.cat((o1, o2), dim=1)
|
||||
mask = torch.sigmoid(mask)
|
||||
|
||||
offset_absmean = torch.mean(torch.abs(offset))
|
||||
if offset_absmean > 50:
|
||||
print(f"Offset abs mean is {offset_absmean}, larger than 50.")
|
||||
|
||||
return torchvision.ops.deform_conv2d(
|
||||
x,
|
||||
offset,
|
||||
self.weight,
|
||||
self.bias,
|
||||
self.stride,
|
||||
self.padding,
|
||||
self.dilation,
|
||||
mask,
|
||||
)
|
||||
|
||||
|
||||
def _no_grad_trunc_normal_(tensor, mean, std, a, b):
|
||||
# From: https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/layers/weight_init.py
|
||||
# Cut & paste from PyTorch official master until it's in a few official releases - RW
|
||||
# Method based on https://people.sc.fsu.edu/~jburkardt/presentations/truncated_normal.pdf
|
||||
def norm_cdf(x):
|
||||
# Computes standard normal cumulative distribution function
|
||||
return (1.0 + math.erf(x / math.sqrt(2.0))) / 2.0
|
||||
|
||||
if (mean < a - 2 * std) or (mean > b + 2 * std):
|
||||
warnings.warn(
|
||||
"mean is more than 2 std from [a, b] in nn.init.trunc_normal_. "
|
||||
"The distribution of values may be incorrect.",
|
||||
stacklevel=2,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
# Values are generated by using a truncated uniform distribution and
|
||||
# then using the inverse CDF for the normal distribution.
|
||||
# Get upper and lower cdf values
|
||||
low = norm_cdf((a - mean) / std)
|
||||
up = norm_cdf((b - mean) / std)
|
||||
|
||||
# Uniformly fill tensor with values from [low, up], then translate to
|
||||
# [2l-1, 2u-1].
|
||||
tensor.uniform_(2 * low - 1, 2 * up - 1)
|
||||
|
||||
# Use inverse cdf transform for normal distribution to get truncated
|
||||
# standard normal
|
||||
tensor.erfinv_()
|
||||
|
||||
# Transform to proper mean, std
|
||||
tensor.mul_(std * math.sqrt(2.0))
|
||||
tensor.add_(mean)
|
||||
|
||||
# Clamp to ensure it's in the proper range
|
||||
tensor.clamp_(min=a, max=b)
|
||||
return tensor
|
||||
|
||||
|
||||
def trunc_normal_(tensor, mean=0.0, std=1.0, a=-2.0, b=2.0):
|
||||
r"""Fills the input Tensor with values drawn from a truncated
|
||||
normal distribution.
|
||||
|
||||
From: https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/layers/weight_init.py
|
||||
|
||||
The values are effectively drawn from the
|
||||
normal distribution :math:`\mathcal{N}(\text{mean}, \text{std}^2)`
|
||||
with values outside :math:`[a, b]` redrawn until they are within
|
||||
the bounds. The method used for generating the random values works
|
||||
best when :math:`a \leq \text{mean} \leq b`.
|
||||
|
||||
Args:
|
||||
tensor: an n-dimensional `torch.Tensor`
|
||||
mean: the mean of the normal distribution
|
||||
std: the standard deviation of the normal distribution
|
||||
a: the minimum cutoff value
|
||||
b: the maximum cutoff value
|
||||
|
||||
Examples:
|
||||
>>> w = torch.empty(3, 5)
|
||||
>>> nn.init.trunc_normal_(w)
|
||||
"""
|
||||
return _no_grad_trunc_normal_(tensor, mean, std, a, b)
|
||||
|
||||
|
||||
# From PyTorch
|
||||
def _ntuple(n):
|
||||
|
||||
def parse(x):
|
||||
if isinstance(x, collections.abc.Iterable):
|
||||
return x
|
||||
return tuple(repeat(x, n))
|
||||
|
||||
return parse
|
||||
|
||||
|
||||
to_1tuple = _ntuple(1)
|
||||
to_2tuple = _ntuple(2)
|
||||
to_3tuple = _ntuple(3)
|
||||
to_4tuple = _ntuple(4)
|
||||
to_ntuple = _ntuple
|
||||
|
||||
|
||||
def closest_larger_multiple_of_minimum_size(size, minimum_size):
|
||||
return int(math.ceil(size / minimum_size) * minimum_size)
|
||||
|
||||
|
||||
class SizeAdapter(object):
|
||||
"""Converts size of input to standard size.
|
||||
Practical deep network works only with input images
|
||||
which height and width are multiples of a minimum size.
|
||||
This class allows to pass to the network images of arbitrary
|
||||
size, by padding the input to the closest multiple
|
||||
and unpadding the network's output to the original size.
|
||||
"""
|
||||
|
||||
def __init__(self, minimum_size=64):
|
||||
self._minimum_size = minimum_size
|
||||
self._pixels_pad_to_width = None
|
||||
self._pixels_pad_to_height = None
|
||||
|
||||
def _closest_larger_multiple_of_minimum_size(self, size):
|
||||
return closest_larger_multiple_of_minimum_size(size, self._minimum_size)
|
||||
|
||||
def pad(self, network_input):
|
||||
"""Returns "network_input" paded with zeros to the "standard" size.
|
||||
The "standard" size correspond to the height and width that
|
||||
are closest multiples of "minimum_size". The method pads
|
||||
height and width and and saves padded values. These
|
||||
values are then used by "unpad_output" method.
|
||||
"""
|
||||
height, width = network_input.size()[-2:]
|
||||
self._pixels_pad_to_height = (
|
||||
self._closest_larger_multiple_of_minimum_size(height) - height
|
||||
)
|
||||
self._pixels_pad_to_width = (
|
||||
self._closest_larger_multiple_of_minimum_size(width) - width
|
||||
)
|
||||
return nn.ZeroPad2d(
|
||||
(self._pixels_pad_to_width, 0, self._pixels_pad_to_height, 0)
|
||||
)(network_input)
|
||||
|
||||
def unpad(self, network_output):
|
||||
"""Returns "network_output" cropped to the original size.
|
||||
The cropping is performed using values save by the "pad_input"
|
||||
method.
|
||||
"""
|
||||
return network_output[
|
||||
..., self._pixels_pad_to_height :, self._pixels_pad_to_width :
|
||||
]
|
||||
|
||||
|
||||
class SmallUpdateBlock(nn.Module):
|
||||
|
||||
def __init__(self, hidden_dim=64, input_dim=64 * 2):
|
||||
super(SmallUpdateBlock, self).__init__()
|
||||
self.gru = ConvGRU(hidden_dim=hidden_dim, input_dim=input_dim)
|
||||
self.res_head = ConvResidualBlocks(num_in_ch=64, num_out_ch=64, num_block=5)
|
||||
|
||||
def forward(self, net, context, motion):
|
||||
|
||||
inp = torch.cat([context, motion], dim=1)
|
||||
net = self.gru(net, inp)
|
||||
delta_net = self.res_head(net)
|
||||
|
||||
return net, delta_net
|
||||
|
||||
|
||||
class ConvGRU(nn.Module):
|
||||
|
||||
def __init__(self, hidden_dim=128, input_dim=192 + 128):
|
||||
super(ConvGRU, self).__init__()
|
||||
self.convz = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
|
||||
self.convr = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
|
||||
self.convq = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
|
||||
|
||||
def forward(self, h, x):
|
||||
hx = torch.cat([h, x], dim=1)
|
||||
|
||||
z = torch.sigmoid(self.convz(hx))
|
||||
r = torch.sigmoid(self.convr(hx))
|
||||
q = torch.tanh(self.convq(torch.cat([r * h, x], dim=1)))
|
||||
|
||||
h = (1 - z) * h + z * q
|
||||
return h
|
||||
|
||||
|
||||
class ConvResidualBlocks(nn.Module):
|
||||
"""Conv and residual block used in BasicVSR.
|
||||
|
||||
Args:
|
||||
num_in_ch (int): Number of input channels. Default: 3.
|
||||
num_out_ch (int): Number of output channels. Default: 64.
|
||||
num_block (int): Number of residual blocks. Default: 15.
|
||||
"""
|
||||
|
||||
def __init__(self, num_in_ch=3, num_out_ch=64, num_block=15):
|
||||
super().__init__()
|
||||
self.main = nn.Sequential(
|
||||
nn.Conv2d(num_in_ch, num_out_ch, 3, 1, 1, bias=True),
|
||||
nn.LeakyReLU(negative_slope=0.1, inplace=True),
|
||||
make_layer(ResidualBlockNoBN, num_block, num_feat=num_out_ch),
|
||||
)
|
||||
|
||||
def forward(self, fea):
|
||||
return self.main(fea)
|
||||
@@ -0,0 +1,170 @@
|
||||
## Files in this folder taken and modified from EvTexture. https://github.com/DachunKai/EvTexture/tree/main/basicsr/archs
|
||||
|
||||
import torch
|
||||
from torch import nn as nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
from .unet_arch import UNet
|
||||
from .arch_util import flow_warp, ConvResidualBlocks, SmallUpdateBlock
|
||||
from .spynet_arch import SpyNet
|
||||
|
||||
|
||||
class EvTexture(nn.Module):
|
||||
"""EvTexture: Event-driven Texture Enhancement for Video Super-Resolution (ICML 2024)
|
||||
Note that: this class is for 4x VSR
|
||||
|
||||
Args:
|
||||
num_feat (int): Number of channels. Default: 64.
|
||||
num_block (int): Number of residual blocks for each branch. Default: 30
|
||||
spynet_path (str): Path to the pretrained weights of SPyNet. Default: None.
|
||||
"""
|
||||
|
||||
def __init__(self, num_feat=64, num_block=30, spynet_path=None):
|
||||
super().__init__()
|
||||
self.num_feat = num_feat
|
||||
|
||||
# RGB-based flowalignment
|
||||
self.spynet = SpyNet(spynet_path)
|
||||
self.cnet = ConvResidualBlocks(num_in_ch=3, num_out_ch=64, num_block=8)
|
||||
|
||||
# iterative texture enhancement module
|
||||
self.enet = UNet(inChannels=1, outChannels=num_feat)
|
||||
self.update_block = SmallUpdateBlock(
|
||||
hidden_dim=num_feat, input_dim=num_feat * 2
|
||||
)
|
||||
self.fusion = nn.Conv2d(num_feat * 2, num_feat, 1, 1, 0, bias=True)
|
||||
|
||||
# propogation
|
||||
self.backward_trunk = ConvResidualBlocks(num_feat + 3, num_feat, num_block)
|
||||
self.forward_trunk = ConvResidualBlocks(num_feat * 2 + 3, num_feat, num_block)
|
||||
|
||||
# reconstruction
|
||||
self.upconv1 = nn.Conv2d(num_feat, num_feat * 4, 3, 1, 1, bias=True)
|
||||
self.upconv2 = nn.Conv2d(num_feat, num_feat * 4, 3, 1, 1, bias=True)
|
||||
self.conv_hr = nn.Conv2d(num_feat, num_feat, 3, 1, 1)
|
||||
self.conv_last = nn.Conv2d(num_feat, 3, 3, 1, 1)
|
||||
|
||||
self.pixel_shuffle = nn.PixelShuffle(2)
|
||||
|
||||
# activation functions
|
||||
self.lrelu = nn.LeakyReLU(negative_slope=0.1, inplace=True)
|
||||
|
||||
def get_flow(self, x):
|
||||
b, n, c, h, w = x.size()
|
||||
|
||||
x_1 = x[:, :-1, :, :, :].reshape(-1, c, h, w)
|
||||
x_2 = x[:, 1:, :, :, :].reshape(-1, c, h, w)
|
||||
|
||||
flows_backward = self.spynet(x_1, x_2).view(b, n - 1, 2, h, w)
|
||||
flows_forward = self.spynet(x_2, x_1).view(b, n - 1, 2, h, w)
|
||||
|
||||
return flows_forward, flows_backward
|
||||
|
||||
# context feature extractor
|
||||
def get_feat(self, x):
|
||||
b, n, c, h, w = x.size()
|
||||
feats_ = self.cnet(x.view(-1, c, h, w))
|
||||
h, w = feats_.shape[2:]
|
||||
feats_ = feats_.view(b, n, -1, h, w)
|
||||
|
||||
return feats_
|
||||
|
||||
def forward(self, imgs, voxels_f, voxels_b):
|
||||
"""Forward function of EvTexture
|
||||
|
||||
Args:
|
||||
imgs: Input frames with shape (b, n, c, h, w). b is batch size. n is the number of frames, and c equals 3 (RGB channels).
|
||||
voxels_f: forward event voxel grids with shape (b, n-1, Bins, h, w). n-1 is intervals between n frames.
|
||||
voxels_b: backward event voxel grids with shape (b, n-1, Bins, h, w).
|
||||
|
||||
Output:
|
||||
out_l: output frames with shape (b, n, c, 4h, 4w)
|
||||
"""
|
||||
|
||||
flows_forward, flows_backward = self.get_flow(imgs)
|
||||
feat_imgs = self.get_feat(imgs)
|
||||
b, n, _, h, w = imgs.size()
|
||||
bins = voxels_f.size()[2]
|
||||
|
||||
# backward branch
|
||||
out_l = []
|
||||
feat_prop = imgs.new_zeros(b, self.num_feat, h, w)
|
||||
for i in range(n - 1, -1, -1):
|
||||
x_i = imgs[:, i, :, :, :]
|
||||
|
||||
if i < n - 1:
|
||||
# motion branch by rgb frames
|
||||
flow = flows_backward[:, i, :, :, :]
|
||||
feat_prop_coarse = flow_warp(feat_prop, flow.permute(0, 2, 3, 1))
|
||||
|
||||
# texture branch by event voxels
|
||||
hidden_state = feat_prop.clone()
|
||||
feat_img = feat_imgs[:, i, :, :, :] # [B, num_feat, H, W]
|
||||
cur_voxel = voxels_f[:, i, :, :, :] # [B, Bins, H, W]
|
||||
|
||||
## iterative update block
|
||||
feat_prop_fine = feat_prop.clone()
|
||||
for j in range(bins - 1, -1, -1):
|
||||
voxel_j = cur_voxel[:, j, :, :].unsqueeze(1) # [B, 1, H, W]
|
||||
feat_motion = self.enet(
|
||||
voxel_j
|
||||
) # [B, num_feat, H, W], enet is UNet(inChannels=1, OurChannels=num_feat)
|
||||
hidden_state, delta_feat = self.update_block(
|
||||
hidden_state, feat_img, feat_motion
|
||||
) # refine coarse hidden state
|
||||
feat_prop_fine = feat_prop_fine + delta_feat
|
||||
|
||||
feat_prop = self.fusion(
|
||||
torch.cat([feat_prop_fine, feat_prop_coarse], dim=1)
|
||||
)
|
||||
|
||||
feat_prop = torch.cat([x_i, feat_prop], dim=1)
|
||||
feat_prop = self.backward_trunk(feat_prop)
|
||||
out_l.insert(0, feat_prop)
|
||||
|
||||
# forward branch
|
||||
feat_prop = torch.zeros_like(feat_prop)
|
||||
for i in range(0, n):
|
||||
x_i = imgs[:, i, :, :, :]
|
||||
|
||||
if i > 0:
|
||||
# motion branch by rgb frames
|
||||
flow = flows_forward[:, i - 1, :, :, :]
|
||||
feat_prop_coarse = flow_warp(feat_prop, flow.permute(0, 2, 3, 1))
|
||||
|
||||
# texture branch by event voxels
|
||||
hidden_state = feat_prop.clone()
|
||||
feat_img = feat_imgs[:, i, :, :, :] # [B, num_feat, H, W]
|
||||
cur_voxel = voxels_b[:, i - 1, :, :, :] # [B, Bins, H, W]
|
||||
|
||||
# iterative update block
|
||||
feat_prop_fine = feat_prop.clone()
|
||||
for j in range(bins - 1, -1, -1):
|
||||
voxel_j = cur_voxel[:, j, :, :].unsqueeze(1) # [B, 1, H, W]
|
||||
feat_motion = self.enet(
|
||||
voxel_j
|
||||
) # [B, num_feat, H, W], enet is UNet(inChannels=1, OurChannels=64)
|
||||
hidden_state, delta_feat = self.update_block(
|
||||
hidden_state, feat_img, feat_motion
|
||||
)
|
||||
feat_prop_fine = feat_prop_fine + delta_feat
|
||||
|
||||
feat_prop = self.fusion(
|
||||
torch.cat([feat_prop_fine, feat_prop_coarse], dim=1)
|
||||
)
|
||||
|
||||
feat_prop = torch.cat([x_i, out_l[i], feat_prop], dim=1)
|
||||
feat_prop = self.forward_trunk(feat_prop)
|
||||
|
||||
# upsample
|
||||
out = self.lrelu(self.pixel_shuffle(self.upconv1(feat_prop)))
|
||||
out = self.lrelu(self.pixel_shuffle(self.upconv2(out)))
|
||||
out = self.lrelu(self.conv_hr(out))
|
||||
out = self.conv_last(out)
|
||||
base = F.interpolate(
|
||||
x_i, scale_factor=4, mode="bilinear", align_corners=False
|
||||
)
|
||||
out += base
|
||||
out_l[i] = out
|
||||
|
||||
return torch.stack(out_l, dim=1)
|
||||
@@ -0,0 +1,160 @@
|
||||
import math
|
||||
import torch
|
||||
from torch import nn as nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
from .arch_util import flow_warp
|
||||
|
||||
|
||||
class BasicModule(nn.Module):
|
||||
"""Basic Module for SpyNet."""
|
||||
|
||||
def __init__(self):
|
||||
super(BasicModule, self).__init__()
|
||||
|
||||
self.basic_module = nn.Sequential(
|
||||
nn.Conv2d(
|
||||
in_channels=8, out_channels=32, kernel_size=7, stride=1, padding=3
|
||||
),
|
||||
nn.ReLU(inplace=False),
|
||||
nn.Conv2d(
|
||||
in_channels=32, out_channels=64, kernel_size=7, stride=1, padding=3
|
||||
),
|
||||
nn.ReLU(inplace=False),
|
||||
nn.Conv2d(
|
||||
in_channels=64, out_channels=32, kernel_size=7, stride=1, padding=3
|
||||
),
|
||||
nn.ReLU(inplace=False),
|
||||
nn.Conv2d(
|
||||
in_channels=32, out_channels=16, kernel_size=7, stride=1, padding=3
|
||||
),
|
||||
nn.ReLU(inplace=False),
|
||||
nn.Conv2d(
|
||||
in_channels=16, out_channels=2, kernel_size=7, stride=1, padding=3
|
||||
),
|
||||
)
|
||||
|
||||
def forward(self, tensor_input):
|
||||
return self.basic_module(tensor_input)
|
||||
|
||||
|
||||
class SpyNet(nn.Module):
|
||||
"""SpyNet architecture.
|
||||
|
||||
Args:
|
||||
load_path (str): path for pretrained SpyNet. Default: None.
|
||||
"""
|
||||
|
||||
def __init__(self, load_path=None):
|
||||
super(SpyNet, self).__init__()
|
||||
self.basic_module = nn.ModuleList([BasicModule() for _ in range(6)])
|
||||
if load_path:
|
||||
self.load_state_dict(
|
||||
torch.load(load_path, map_location=lambda storage, loc: storage)[
|
||||
"params"
|
||||
]
|
||||
)
|
||||
|
||||
self.register_buffer(
|
||||
"mean", torch.Tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
|
||||
)
|
||||
self.register_buffer(
|
||||
"std", torch.Tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)
|
||||
)
|
||||
|
||||
def preprocess(self, tensor_input):
|
||||
tensor_output = (tensor_input - self.mean) / self.std
|
||||
return tensor_output
|
||||
|
||||
def process(self, ref, supp):
|
||||
flow = []
|
||||
|
||||
ref = [self.preprocess(ref)]
|
||||
supp = [self.preprocess(supp)]
|
||||
|
||||
for level in range(5):
|
||||
ref.insert(
|
||||
0,
|
||||
F.avg_pool2d(
|
||||
input=ref[0], kernel_size=2, stride=2, count_include_pad=False
|
||||
),
|
||||
)
|
||||
supp.insert(
|
||||
0,
|
||||
F.avg_pool2d(
|
||||
input=supp[0], kernel_size=2, stride=2, count_include_pad=False
|
||||
),
|
||||
)
|
||||
|
||||
flow = ref[0].new_zeros(
|
||||
[
|
||||
ref[0].size(0),
|
||||
2,
|
||||
int(math.floor(ref[0].size(2) / 2.0)),
|
||||
int(math.floor(ref[0].size(3) / 2.0)),
|
||||
]
|
||||
)
|
||||
|
||||
for level in range(len(ref)):
|
||||
upsampled_flow = (
|
||||
F.interpolate(
|
||||
input=flow, scale_factor=2, mode="bilinear", align_corners=True
|
||||
)
|
||||
* 2.0
|
||||
)
|
||||
|
||||
if upsampled_flow.size(2) != ref[level].size(2):
|
||||
upsampled_flow = F.pad(
|
||||
input=upsampled_flow, pad=[0, 0, 0, 1], mode="replicate"
|
||||
)
|
||||
if upsampled_flow.size(3) != ref[level].size(3):
|
||||
upsampled_flow = F.pad(
|
||||
input=upsampled_flow, pad=[0, 1, 0, 0], mode="replicate"
|
||||
)
|
||||
|
||||
flow = (
|
||||
self.basic_module[level](
|
||||
torch.cat(
|
||||
[
|
||||
ref[level],
|
||||
flow_warp(
|
||||
supp[level],
|
||||
upsampled_flow.permute(0, 2, 3, 1),
|
||||
interp_mode="bilinear",
|
||||
padding_mode="border",
|
||||
),
|
||||
upsampled_flow,
|
||||
],
|
||||
1,
|
||||
)
|
||||
)
|
||||
+ upsampled_flow
|
||||
)
|
||||
|
||||
return flow
|
||||
|
||||
def forward(self, ref, supp):
|
||||
assert ref.size() == supp.size()
|
||||
|
||||
h, w = ref.size(2), ref.size(3)
|
||||
w_floor = math.floor(math.ceil(w / 32.0) * 32.0)
|
||||
h_floor = math.floor(math.ceil(h / 32.0) * 32.0)
|
||||
|
||||
ref = F.interpolate(
|
||||
input=ref, size=(h_floor, w_floor), mode="bilinear", align_corners=False
|
||||
)
|
||||
supp = F.interpolate(
|
||||
input=supp, size=(h_floor, w_floor), mode="bilinear", align_corners=False
|
||||
)
|
||||
|
||||
flow = F.interpolate(
|
||||
input=self.process(ref, supp),
|
||||
size=(h, w),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
)
|
||||
|
||||
flow[:, 0, :, :] *= float(w) / float(w_floor)
|
||||
flow[:, 1, :, :] *= float(h) / float(h_floor)
|
||||
|
||||
return flow
|
||||
@@ -0,0 +1,130 @@
|
||||
## Modified from timelens. https://github.com/uzh-rpg/rpg_timelens/blob/main/timelens/superslomo/unet.py
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from .arch_util import SizeAdapter
|
||||
from torch import nn
|
||||
|
||||
|
||||
class up(nn.Module):
|
||||
def __init__(self, inChannels, outChannels):
|
||||
super(up, self).__init__()
|
||||
self.conv1 = nn.Conv2d(inChannels, outChannels, 3, stride=1, padding=1)
|
||||
self.conv2 = nn.Conv2d(2 * outChannels, outChannels, 3, stride=1, padding=1)
|
||||
|
||||
def forward(self, x, skpCn):
|
||||
x = F.interpolate(x, scale_factor=2, mode="bilinear")
|
||||
x = F.leaky_relu(self.conv1(x), negative_slope=0.1)
|
||||
x = F.leaky_relu(self.conv2(torch.cat((x, skpCn), 1)), negative_slope=0.1)
|
||||
return x
|
||||
|
||||
|
||||
class down(nn.Module):
|
||||
def __init__(self, inChannels, outChannels, filterSize):
|
||||
super(down, self).__init__()
|
||||
self.conv1 = nn.Conv2d(
|
||||
inChannels,
|
||||
outChannels,
|
||||
filterSize,
|
||||
stride=1,
|
||||
padding=int((filterSize - 1) / 2),
|
||||
)
|
||||
self.conv2 = nn.Conv2d(
|
||||
outChannels,
|
||||
outChannels,
|
||||
filterSize,
|
||||
stride=1,
|
||||
padding=int((filterSize - 1) / 2),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.avg_pool2d(x, 2)
|
||||
x = F.leaky_relu(self.conv1(x), negative_slope=0.1)
|
||||
x = F.leaky_relu(self.conv2(x), negative_slope=0.1)
|
||||
return x
|
||||
|
||||
|
||||
class UNet(nn.Module):
|
||||
"""Modified version of Unet from SuperSloMo.
|
||||
|
||||
Difference :
|
||||
1) there is an option to skip ReLU after the last convolution.
|
||||
2) there is a size adapter module that makes sure that input of all sizes
|
||||
can be processed correctly. It is necessary because original
|
||||
UNet can process only inputs with spatial dimensions divisible by 32.
|
||||
"""
|
||||
|
||||
def __init__(self, inChannels, outChannels, ends_with_relu=True, load_path=None):
|
||||
super(UNet, self).__init__()
|
||||
self._ends_with_relu = ends_with_relu
|
||||
self._size_adapter = SizeAdapter(minimum_size=32)
|
||||
|
||||
# 5-level
|
||||
self.conv1 = nn.Conv2d(inChannels, 8, 7, stride=1, padding=3)
|
||||
self.conv2 = nn.Conv2d(8, 8, 7, stride=1, padding=3)
|
||||
self.down1 = down(8, 16, 5)
|
||||
self.down2 = down(16, 32, 3)
|
||||
self.down3 = down(32, 64, 3)
|
||||
self.down4 = down(64, 128, 3)
|
||||
self.down5 = down(128, 128, 3)
|
||||
self.up1 = up(128, 128)
|
||||
self.up2 = up(128, 64)
|
||||
self.up3 = up(64, 32)
|
||||
self.up4 = up(32, 16)
|
||||
self.up5 = up(16, 8)
|
||||
self.conv3 = nn.Conv2d(8, outChannels, 3, stride=1, padding=1)
|
||||
|
||||
if load_path:
|
||||
self.load_state_dict(
|
||||
torch.load(load_path, map_location=lambda storage, loc: storage)[
|
||||
"params_ema"
|
||||
]
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self._size_adapter.pad(x)
|
||||
x = F.leaky_relu(self.conv1(x), negative_slope=0.1)
|
||||
s1 = F.leaky_relu(self.conv2(x), negative_slope=0.1)
|
||||
s2 = self.down1(s1)
|
||||
s3 = self.down2(s2)
|
||||
s4 = self.down3(s3)
|
||||
s5 = self.down4(s4)
|
||||
x = self.down5(s5)
|
||||
x = self.up1(x, s5)
|
||||
x = self.up2(x, s4)
|
||||
x = self.up3(x, s3)
|
||||
x = self.up4(x, s2)
|
||||
x = self.up5(x, s1)
|
||||
|
||||
# Note that original code has relu et the end.
|
||||
if self._ends_with_relu == True:
|
||||
x = F.leaky_relu(self.conv3(x), negative_slope=0.1)
|
||||
else:
|
||||
x = self.conv3(x)
|
||||
# Size adapter crops the output to the original size.
|
||||
x = self._size_adapter.unpad(x)
|
||||
return x
|
||||
|
||||
|
||||
def patch_chunk_2x(input):
|
||||
"""
|
||||
input (Tensor): [B, C, H, W], and H, W are divisible by 2.
|
||||
|
||||
return:
|
||||
result (Tensor): [B, 4C, H/2, H/W]
|
||||
"""
|
||||
result = []
|
||||
split_h = torch.chunk(input, 2, -2)
|
||||
for sli in split_h:
|
||||
sli_w = torch.chunk(sli, 2, -1)
|
||||
for i in range(2):
|
||||
result.append(sli_w[i])
|
||||
assert len(result) == 4
|
||||
result = torch.cat(result, dim=1)
|
||||
return result
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
net = UNet(1, 2)
|
||||
input = torch.randn((4, 1, 64, 64))
|
||||
out = net(input)
|
||||
@@ -0,0 +1,14 @@
|
||||
[project]
|
||||
name = "comfyui-evtexture"
|
||||
description = "Wrapper for EvTexture Video Upscaler: [a/https://github.com/DachunKai/EvTexture](https://github.com/DachunKai/EvTexture)"
|
||||
version = "1.0.0"
|
||||
dependencies = ["torch", "torchvision>=0.9.0"]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/tocubed/ComfyUI-EvTexture"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "tocubed"
|
||||
DisplayName = "ComfyUI-EvTexture"
|
||||
Icon = ""
|
||||
@@ -0,0 +1,2 @@
|
||||
torch
|
||||
torchvision>=0.9.0
|
||||
Reference in New Issue
Block a user