commit e8613cb3b361a82e0bd9d6626b3a4105399554b7 Author: Priyank Patel Date: Wed Jan 1 18:48:42 2025 -0800 Init diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..ed8ebf5 --- /dev/null +++ b/.gitignore @@ -0,0 +1 @@ +__pycache__ \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..f9da786 --- /dev/null +++ b/__init__.py @@ -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", +} diff --git a/esim.py b/esim.py new file mode 100644 index 0000000..0b27f44 --- /dev/null +++ b/esim.py @@ -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 diff --git a/evoxels.py b/evoxels.py new file mode 100644 index 0000000..ffb6af6 --- /dev/null +++ b/evoxels.py @@ -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 diff --git a/evtexture/arch_util.py b/evtexture/arch_util.py new file mode 100644 index 0000000..60da562 --- /dev/null +++ b/evtexture/arch_util.py @@ -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) diff --git a/evtexture/evtexture_arch.py b/evtexture/evtexture_arch.py new file mode 100644 index 0000000..a867c27 --- /dev/null +++ b/evtexture/evtexture_arch.py @@ -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) diff --git a/evtexture/spynet_arch.py b/evtexture/spynet_arch.py new file mode 100644 index 0000000..5c8e608 --- /dev/null +++ b/evtexture/spynet_arch.py @@ -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 diff --git a/evtexture/unet_arch.py b/evtexture/unet_arch.py new file mode 100644 index 0000000..6593aa0 --- /dev/null +++ b/evtexture/unet_arch.py @@ -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) diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..6481ea8 --- /dev/null +++ b/pyproject.toml @@ -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 = "" \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..4453ec4 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,2 @@ +torch +torchvision>=0.9.0 \ No newline at end of file