This commit is contained in:
Priyank Patel
2025-01-01 18:48:42 -08:00
commit e8613cb3b3
10 changed files with 1420 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
__pycache__
+144
View File
@@ -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",
}
+98
View File
@@ -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
View File
@@ -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
+439
View File
@@ -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)
+170
View File
@@ -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)
+160
View File
@@ -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
+130
View File
@@ -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)
+14
View File
@@ -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 = ""
+2
View File
@@ -0,0 +1,2 @@
torch
torchvision>=0.9.0