Fix import errors when using an alternative custom nodes path

Done by inserting the repository directory into sys.path in '__init__.py' and modifying all import statements to remove `custom_nodes.facerestore_cf.`
This commit is contained in:
Quaggles
2024-10-13 02:16:28 +10:00
parent 67f90bc6be
commit 589ab5d629
31 changed files with 73 additions and 71 deletions
+6 -4
View File
@@ -1,17 +1,19 @@
import os
import sys
repo_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.insert(0, repo_dir)
from comfy import model_management
import torch
import comfy.utils
import numpy as np
import cv2
import math
from custom_nodes.facerestore_cf.facelib.utils.face_restoration_helper import FaceRestoreHelper
from custom_nodes.facerestore_cf.facelib.detection.retinaface import retinaface
from facelib.utils.face_restoration_helper import FaceRestoreHelper
from facelib.detection.retinaface import retinaface
from torchvision.transforms.functional import normalize
from comfy_extras.chainner_models import model_loading
import folder_paths
import sys
from custom_nodes.facerestore_cf.basicsr.utils.registry import ARCH_REGISTRY
from basicsr.utils.registry import ARCH_REGISTRY
# import codeformer_arch
dir_facerestore_models = os.path.join(folder_paths.models_dir, "facerestore_models")
+3 -3
View File
@@ -2,8 +2,8 @@ import importlib
from copy import deepcopy
from os import path as osp
from custom_nodes.facerestore_cf.basicsr.utils import get_root_logger, scandir
from custom_nodes.facerestore_cf.basicsr.utils.registry import ARCH_REGISTRY
from basicsr.utils import get_root_logger, scandir
from basicsr.utils.registry import ARCH_REGISTRY
__all__ = ['build_network']
@@ -13,7 +13,7 @@ __all__ = ['build_network']
arch_folder = osp.dirname(osp.abspath(__file__))
arch_filenames = [osp.splitext(osp.basename(v))[0] for v in scandir(arch_folder) if v.endswith('_arch.py')]
# import all the arch modules
_arch_modules = [importlib.import_module(f'custom_nodes.facerestore_cf.basicsr.archs.{file_name}') for file_name in arch_filenames]
_arch_modules = [importlib.import_module(f'basicsr.archs.{file_name}') for file_name in arch_filenames]
def build_network(opt):
+1 -1
View File
@@ -1,5 +1,5 @@
import torch.nn as nn
from custom_nodes.facerestore_cf.basicsr.utils.registry import ARCH_REGISTRY
from basicsr.utils.registry import ARCH_REGISTRY
def conv3x3(inplanes, outplanes, stride=1):
+2 -2
View File
@@ -10,8 +10,8 @@ from torch.nn import functional as F
from torch.nn import init as init
from torch.nn.modules.batchnorm import _BatchNorm
from custom_nodes.facerestore_cf.basicsr.ops.dcn import ModulatedDeformConvPack, modulated_deform_conv
from custom_nodes.facerestore_cf.basicsr.utils import get_root_logger
from basicsr.ops.dcn import ModulatedDeformConvPack, modulated_deform_conv
from basicsr.utils import get_root_logger
@torch.no_grad()
+3 -3
View File
@@ -5,9 +5,9 @@ from torch import nn, Tensor
import torch.nn.functional as F
from typing import Optional, List
from custom_nodes.facerestore_cf.basicsr.archs.vqgan_arch import *
from custom_nodes.facerestore_cf.basicsr.utils import get_root_logger
from custom_nodes.facerestore_cf.basicsr.utils.registry import ARCH_REGISTRY
from basicsr.archs.vqgan_arch import *
from basicsr.utils import get_root_logger
from basicsr.utils.registry import ARCH_REGISTRY
def calc_mean_std(feat, eps=1e-5):
"""Calculate mean and std for adaptive_instance_normalization.
+1 -1
View File
@@ -2,7 +2,7 @@ import torch
from torch import nn as nn
from torch.nn import functional as F
from custom_nodes.facerestore_cf.basicsr.utils.registry import ARCH_REGISTRY
from basicsr.utils.registry import ARCH_REGISTRY
from .arch_util import default_init_weights, make_layer, pixel_unshuffle
+1 -1
View File
@@ -4,7 +4,7 @@ from collections import OrderedDict
from torch import nn as nn
from torchvision.models import vgg as vgg
from custom_nodes.facerestore_cf.basicsr.utils.registry import ARCH_REGISTRY
from basicsr.utils.registry import ARCH_REGISTRY
VGG_PRETRAIN_PATH = 'experiments/pretrained_models/vgg19-dcbb9e9d.pth'
NAMES = {
+2 -2
View File
@@ -8,8 +8,8 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
import copy
from custom_nodes.facerestore_cf.basicsr.utils import get_root_logger
from custom_nodes.facerestore_cf.basicsr.utils.registry import ARCH_REGISTRY
from basicsr.utils import get_root_logger
from basicsr.utils.registry import ARCH_REGISTRY
def normalize(in_channels):
return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
+4 -4
View File
@@ -7,10 +7,10 @@ from copy import deepcopy
from functools import partial
from os import path as osp
from custom_nodes.facerestore_cf.basicsr.data.prefetch_dataloader import PrefetchDataLoader
from custom_nodes.facerestore_cf.basicsr.utils import get_root_logger, scandir
from custom_nodes.facerestore_cf.basicsr.utils.dist_util import get_dist_info
from custom_nodes.facerestore_cf.basicsr.utils.registry import DATASET_REGISTRY
from basicsr.data.prefetch_dataloader import PrefetchDataLoader
from basicsr.utils import get_root_logger, scandir
from basicsr.utils.dist_util import get_dist_info
from basicsr.utils.registry import DATASET_REGISTRY
__all__ = ['build_dataset', 'build_dataloader']
+2 -2
View File
@@ -4,8 +4,8 @@ import torch
from os import path as osp
from torch.nn import functional as F
from custom_nodes.facerestore_cf.basicsr.data.transforms import mod_crop
from custom_nodes.facerestore_cf.basicsr.utils import img2tensor, scandir
from basicsr.data.transforms import mod_crop
from basicsr.utils import img2tensor, scandir
def read_img_seq(path, require_mod_crop=False, scale=1):
+2 -2
View File
@@ -1,7 +1,7 @@
from copy import deepcopy
from custom_nodes.facerestore_cf.basicsr.utils import get_root_logger
from custom_nodes.facerestore_cf.basicsr.utils.registry import LOSS_REGISTRY
from basicsr.utils import get_root_logger
from basicsr.utils.registry import LOSS_REGISTRY
from .losses import (CharbonnierLoss, GANLoss, L1Loss, MSELoss, PerceptualLoss, WeightedTVLoss, g_path_regularize,
gradient_penalty_loss, r1_penalty)
+2 -2
View File
@@ -5,8 +5,8 @@ from torch import autograd as autograd
from torch import nn as nn
from torch.nn import functional as F
from custom_nodes.facerestore_cf.basicsr.archs.vgg_arch import VGGFeatureExtractor
from custom_nodes.facerestore_cf.basicsr.utils.registry import LOSS_REGISTRY
from basicsr.archs.vgg_arch import VGGFeatureExtractor
from basicsr.utils.registry import LOSS_REGISTRY
from .loss_util import weighted_loss
_reduction_modes = ['none', 'mean', 'sum']
+1 -1
View File
@@ -1,6 +1,6 @@
from copy import deepcopy
from custom_nodes.facerestore_cf.basicsr.utils.registry import METRIC_REGISTRY
from basicsr.utils.registry import METRIC_REGISTRY
from .psnr_ssim import calculate_psnr, calculate_ssim
__all__ = ['calculate_psnr', 'calculate_ssim']
+1 -1
View File
@@ -1,6 +1,6 @@
import numpy as np
from custom_nodes.facerestore_cf.basicsr.utils.matlab_functions import bgr2ycbcr
from basicsr.utils.matlab_functions import bgr2ycbcr
def reorder_image(img, input_order='HWC'):
+2 -2
View File
@@ -1,8 +1,8 @@
import cv2
import numpy as np
from custom_nodes.facerestore_cf.basicsr.metrics.metric_util import reorder_image, to_y_channel
from custom_nodes.facerestore_cf.basicsr.utils.registry import METRIC_REGISTRY
from basicsr.metrics.metric_util import reorder_image, to_y_channel
from basicsr.utils.registry import METRIC_REGISTRY
@METRIC_REGISTRY.register()
+2 -2
View File
@@ -2,8 +2,8 @@ import importlib
from copy import deepcopy
from os import path as osp
from custom_nodes.facerestore_cf.basicsr.utils import get_root_logger, scandir
from custom_nodes.facerestore_cf.basicsr.utils.registry import MODEL_REGISTRY
from basicsr.utils import get_root_logger, scandir
from basicsr.utils.registry import MODEL_REGISTRY
__all__ = ['build_model']
+7 -7
View File
@@ -8,14 +8,14 @@ import time
import torch
from os import path as osp
from custom_nodes.facerestore_cf.basicsr.data import build_dataloader, build_dataset
from custom_nodes.facerestore_cf.basicsr.data.data_sampler import EnlargedSampler
from custom_nodes.facerestore_cf.basicsr.data.prefetch_dataloader import CPUPrefetcher, CUDAPrefetcher
from custom_nodes.facerestore_cf.basicsr.models import build_model
from custom_nodes.facerestore_cf.basicsr.utils import (MessageLogger, check_resume, get_env_info, get_root_logger, init_tb_logger,
from basicsr.data import build_dataloader, build_dataset
from basicsr.data.data_sampler import EnlargedSampler
from basicsr.data.prefetch_dataloader import CPUPrefetcher, CUDAPrefetcher
from basicsr.models import build_model
from basicsr.utils import (MessageLogger, check_resume, get_env_info, get_root_logger, init_tb_logger,
init_wandb_logger, make_exp_dirs, mkdir_and_rename, set_random_seed)
from custom_nodes.facerestore_cf.basicsr.utils.dist_util import get_dist_info, init_dist
from custom_nodes.facerestore_cf.basicsr.utils.options import dict2str, parse
from basicsr.utils.dist_util import get_dist_info, init_dist
from basicsr.utils.options import dict2str, parse
import warnings
# ignore UserWarning: Detected call of `lr_scheduler.step()` before `optimizer.step()`.
+1 -1
View File
@@ -149,7 +149,7 @@ def get_env_info():
import torch
import torchvision
from custom_nodes.facerestore_cf.basicsr.version import __version__
from basicsr.version import __version__
msg = r"""
____ _ _____ ____
/ __ ) ____ _ _____ (_)_____/ ___/ / __ \
+1 -1
View File
@@ -2,7 +2,7 @@ import yaml
import time
from collections import OrderedDict
from os import path as osp
from custom_nodes.facerestore_cf.basicsr.utils.misc import get_time_str
from basicsr.utils.misc import get_time_str
def ordered_yaml():
"""Support OrderedDict for yaml.
+1 -1
View File
@@ -5,7 +5,7 @@ import os
import queue
import threading
import torch
from custom_nodes.facerestore_cf.basicsr.utils.download_util import load_file_from_url
from basicsr.utils.download_util import load_file_from_url
from torch.nn import functional as F
# ROOT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+3 -3
View File
@@ -5,9 +5,9 @@ from torch import nn, Tensor
import torch.nn.functional as F
from typing import Optional, List
from custom_nodes.facerestore_cf.basicsr.archs.vqgan_arch import *
from custom_nodes.facerestore_cf.basicsr.utils import get_root_logger
from custom_nodes.facerestore_cf.basicsr.utils.registry import ARCH_REGISTRY
from basicsr.archs.vqgan_arch import *
from basicsr.utils import get_root_logger
from basicsr.utils.registry import ARCH_REGISTRY
def calc_mean_std(feat, eps=1e-5):
"""Calculate mean and std for adaptive_instance_normalization.
+3 -3
View File
@@ -4,9 +4,9 @@ from torch import nn
from copy import deepcopy
import pathlib
from custom_nodes.facerestore_cf.facelib.utils import load_file_from_url
from custom_nodes.facerestore_cf.facelib.utils import download_pretrained_models
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.models.common import Conv
from facelib.utils import load_file_from_url
from facelib.utils import download_pretrained_models
from facelib.detection.yolov5face.models.common import Conv
from .retinaface.retinaface import RetinaFace
from .yolov5face.face_detector import YoloDetector
+3 -3
View File
@@ -6,9 +6,9 @@ import torch.nn.functional as F
from PIL import Image
from torchvision.models._utils import IntermediateLayerGetter as IntermediateLayerGetter
from comfy import model_management
from custom_nodes.facerestore_cf.facelib.detection.align_trans import get_reference_facial_points, warp_and_crop_face
from custom_nodes.facerestore_cf.facelib.detection.retinaface.retinaface_net import FPN, SSH, MobileNetV1, make_bbox_head, make_class_head, make_landmark_head
from custom_nodes.facerestore_cf.facelib.detection.retinaface.retinaface_utils import (PriorBox, batched_decode, batched_decode_landm, decode, decode_landm,
from facelib.detection.align_trans import get_reference_facial_points, warp_and_crop_face
from facelib.detection.retinaface.retinaface_net import FPN, SSH, MobileNetV1, make_bbox_head, make_class_head, make_landmark_head
from facelib.detection.retinaface.retinaface_utils import (PriorBox, batched_decode, batched_decode_landm, decode, decode_landm,
py_cpu_nms)
# device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
@@ -7,10 +7,10 @@ import numpy as np
import torch
from torch import nn
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.models.common import Conv
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.models.yolo import Model
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.utils.datasets import letterbox
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.utils.general import (
from facelib.detection.yolov5face.models.common import Conv
from facelib.detection.yolov5face.models.yolo import Model
from facelib.detection.yolov5face.utils.datasets import letterbox
from facelib.detection.yolov5face.utils.general import (
check_img_size,
non_max_suppression_face,
scale_coords,
@@ -6,8 +6,8 @@ import numpy as np
import torch
from torch import nn
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.utils.datasets import letterbox
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.utils.general import (
from facelib.detection.yolov5face.utils.datasets import letterbox
from facelib.detection.yolov5face.utils.general import (
make_divisible,
non_max_suppression,
scale_coords,
@@ -4,7 +4,7 @@ import numpy as np
import torch
from torch import nn
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.models.common import Conv
from facelib.detection.yolov5face.models.common import Conv
class CrossConv(nn.Module):
+5 -5
View File
@@ -6,7 +6,7 @@ import torch
import yaml # for torch hub
from torch import nn
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.models.common import (
from facelib.detection.yolov5face.models.common import (
C3,
NMS,
SPP,
@@ -20,10 +20,10 @@ from custom_nodes.facerestore_cf.facelib.detection.yolov5face.models.common impo
ShuffleV2Block,
StemBlock,
)
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.models.experimental import CrossConv, MixConv2d
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.utils.autoanchor import check_anchor_order
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.utils.general import make_divisible
from custom_nodes.facerestore_cf.facelib.detection.yolov5face.utils.torch_utils import copy_attr, fuse_conv_and_bn
from facelib.detection.yolov5face.models.experimental import CrossConv, MixConv2d
from facelib.detection.yolov5face.utils.autoanchor import check_anchor_order
from facelib.detection.yolov5face.utils.general import make_divisible
from facelib.detection.yolov5face.utils.torch_utils import copy_attr, fuse_conv_and_bn
class Detect(nn.Module):
+1 -1
View File
@@ -1,6 +1,6 @@
import torch
from custom_nodes.facerestore_cf.facelib.utils import load_file_from_url
from facelib.utils import load_file_from_url
from .bisenet import BiSeNet
from .parsenet import ParseNet
+3 -3
View File
@@ -4,9 +4,9 @@ import os
import torch
from torchvision.transforms.functional import normalize
from custom_nodes.facerestore_cf.facelib.detection import init_detection_model
from custom_nodes.facerestore_cf.facelib.parsing import init_parsing_model
from custom_nodes.facerestore_cf.facelib.utils.misc import img2tensor, imwrite
from facelib.detection import init_detection_model
from facelib.parsing import init_parsing_model
from facelib.utils.misc import img2tensor, imwrite
def get_largest_face(det_faces, h, w):
+2 -2
View File
@@ -211,8 +211,8 @@ def paste_face_back(img, face, inverse_affine):
if __name__ == '__main__':
import os
from custom_nodes.facerestore_cf.facelib.detection import init_detection_model
from custom_nodes.facerestore_cf.facelib.utils.face_restoration_helper import get_largest_face
from facelib.detection import init_detection_model
from facelib.utils.face_restoration_helper import get_largest_face
img_path = '/home/wxt/datasets/ffhq/ffhq_wild/00009.png'
img_name = os.splitext(os.path.basename(img_path))[0]
+1 -1
View File
@@ -6,7 +6,7 @@ from torch.hub import download_url_to_file, get_dir
from urllib.parse import urlparse
import folder_paths
# from custom_nodes.facerestore_cf.basicsr.utils.download_util import download_file_from_google_drive
# from basicsr.utils.download_util import download_file_from_google_drive
#import gdown