Fix import

This commit is contained in:
peteromallet
2024-05-25 21:24:16 +02:00
parent 8370a2f1ef
commit b9f4bf6730
4 changed files with 14 additions and 52 deletions
@@ -1,38 +0,0 @@
import os
import sys
sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
import shutil
import torch
import torch.nn.functional as F
import PIL
import torchvision.transforms.functional as transform
from vfi_utils import load_file_from_github_release
from vfi_models import gmfss_fortuna, ifrnet, ifunet, m2m, rife, sepconv, amt, xvfi, cain, flavr
import numpy as np
frame_0 = torch.from_numpy(np.array(PIL.Image.open("demo_frames/anime0.png").convert("RGB")).astype(np.float32) / 255.0).unsqueeze(0)
frame_1 = torch.from_numpy(np.array(PIL.Image.open("demo_frames/anime1.png").convert("RGB")).astype(np.float32) / 255.0).unsqueeze(0)
if os.path.exists("test_result"):
shutil.rmtree("test_result")
vfi_node_class = gmfss_fortuna.GMFSS_Fortuna_VFI()
for i, ckpt_name in enumerate(vfi_node_class.INPUT_TYPES()["required"]["ckpt_name"][0][:2]):
result = vfi_node_class.vfi(ckpt_name, torch.cat([
frame_0,
frame_1,
frame_0,
frame_1
], dim=0).cuda(), multipler=4, batch_size=2)[0]
print(result.shape)
print(f"Generated {result.size(0)} frames")
frames = [PIL.Image.fromarray(np.clip((frame * 255).numpy(), 0, 255).astype(np.uint8)) for frame in result]
print(result[0].shape)
os.makedirs(f"test_result/video{i}", exist_ok=True)
for j, frame in enumerate(frames):
frame.save(f"test_result/video{i}/{j}.jpg")
frames[0].save(f"test_result/video{i}.gif", save_all=True, append_images=frames[1:], optimize=True, duration=1/3, loop=0)
os.startfile(f"test_result{os.path.sep}video{i}.gif")
#torchvision.io.video.write_video("test.mp4", einops.rearrange(result, "n c h w -> n h w c").cpu(), fps=1)
@@ -3,7 +3,7 @@ from comfy.model_management import get_torch_device, soft_empty_cache
import bisect
import numpy as np
import typing
from vfi_utils import InterpolationStateListImport, load_file_from_github_release, preprocess_frames, postprocess_frames
from import_vfi_utils import InterpolationStateListImport, load_file_from_github_release, preprocess_frames, postprocess_frames
import pathlib
import gc
@@ -80,7 +80,7 @@ from torch import nn
from torch.nn import functional as F
class SubTreeExtractor(nn.Module):
class SubTreeExtractorImport(nn.Module):
"""Extracts a hierarchical set of features from an image.
This is a conventional, hierarchical image feature extractor, that extracts
@@ -121,13 +121,13 @@ class SubTreeExtractor(nn.Module):
return pyramid
class FeatureExtractor(nn.Module):
class FeatureExtractorImport(nn.Module):
"""Extracts features from an image pyramid using a cascaded architecture.
"""
def __init__(self, in_channels=3, channels=64, sub_levels=4):
super().__init__()
self.extract_sublevels = SubTreeExtractor(in_channels, channels, sub_levels)
self.extract_sublevels = SubTreeExtractorImport(in_channels, channels, sub_levels)
self.sub_levels = sub_levels
def forward(self, image_pyramid: List[torch.Tensor]) -> List[torch.Tensor]:
@@ -216,7 +216,7 @@ def get_channels_at_level(level, filters):
return (sum(filters << i for i in range(level)) + channels + flows) * n_images
class Fusion(nn.Module):
class FusionImport(nn.Module):
"""The decoder."""
def __init__(self, n_layers=4, specialized_layers=3, filters=64):
@@ -373,7 +373,7 @@ from torch import nn
class Interpolator(nn.Module):
class InterpolatorImport(nn.Module):
def __init__(
self,
pyramid_levels=7,
@@ -388,9 +388,9 @@ class Interpolator(nn.Module):
self.pyramid_levels = pyramid_levels
self.fusion_pyramid_levels = fusion_pyramid_levels
self.extract = FeatureExtractor(3, filters, sub_levels)
self.predict_flow = PyramidFlowEstimator(filters, flow_convs, flow_filters)
self.fuse = Fusion(sub_levels, specialized_levels, filters)
self.extract = FeatureExtractorImport(3, filters, sub_levels)
self.predict_flow = PyramidFlowEstimatorImport(filters, flow_convs, flow_filters)
self.fuse = FusionImport(sub_levels, specialized_levels, filters)
def shuffle_images(self, x0, x1):
return [
@@ -497,7 +497,7 @@ from torch.nn import functional as F
class FlowEstimator(nn.Module):
class FlowEstimatorImport(nn.Module):
"""Small-receptive field predictor for computing the flow between two images.
This is used to compute the residual flow fields in PyramidFlowEstimator.
@@ -513,7 +513,7 @@ class FlowEstimator(nn.Module):
"""
def __init__(self, in_channels: int, num_convs: int, num_filters: int):
super(FlowEstimator, self).__init__()
super(FlowEstimatorImport, self).__init__()
self._convs = nn.ModuleList()
for i in range(num_convs):
@@ -543,20 +543,20 @@ class FlowEstimator(nn.Module):
return net
class PyramidFlowEstimator(nn.Module):
class PyramidFlowEstimatorImport(nn.Module):
"""Predicts optical flow by coarse-to-fine refinement.
"""
def __init__(self, filters: int = 64,
flow_convs: tuple = (3, 3, 3, 3),
flow_filters: tuple = (32, 64, 128, 256)):
super(PyramidFlowEstimator, self).__init__()
super(PyramidFlowEstimatorImport, self).__init__()
in_channels = filters << 1
predictors = []
for i in range(len(flow_convs)):
predictors.append(
FlowEstimator(
FlowEstimatorImport(
in_channels=in_channels,
num_convs=flow_convs[i],
num_filters=flow_filters[i]))