# !/usr/bin/env python # -*- coding: UTF-8 -*- import os import torch import gc from PIL import Image import numpy as np import cv2 from comfy.utils import common_upscale,ProgressBar cur_path = os.path.dirname(os.path.abspath(__file__)) def gc_cleanup(): gc.collect() torch.cuda.empty_cache() def tensor2cv(tensor_image): if len(tensor_image.shape)==4:# b hwc to hwc tensor_image=tensor_image.squeeze(0) if tensor_image.is_cuda: tensor_image = tensor_image.cpu() tensor_image=tensor_image.numpy() #εε½’δΈ€εŒ– maxValue=tensor_image.max() tensor_image=tensor_image*255/maxValue img_cv2=np.uint8(tensor_image)#32 to uint8 img_cv2=cv2.cvtColor(img_cv2,cv2.COLOR_RGB2BGR) return img_cv2 def phi2narry(img): img = torch.from_numpy(np.array(img).astype(np.float32) / 255.0).unsqueeze(0) return img def tensor2image(tensor): tensor = tensor.cpu() image_np = tensor.squeeze().mul(255).clamp(0, 255).byte().numpy() image = Image.fromarray(image_np, mode='RGB') return image def tensor2pillist(tensor_in): d1, _, _, _ = tensor_in.size() if d1 == 1: img_list = [tensor2image(tensor_in)] else: tensor_list = torch.chunk(tensor_in, chunks=d1) img_list=[tensor2image(i) for i in tensor_list] return img_list def tensor2pillist_upscale(tensor_in,width,height): d1, _, _, _ = tensor_in.size() if d1 == 1: img_list = [nomarl_upscale(tensor_in,width,height)] else: tensor_list = torch.chunk(tensor_in, chunks=d1) img_list=[nomarl_upscale(i,width,height) for i in tensor_list] return img_list def tensor2list(tensor_in,width,height): if tensor_in is None: return None d1, _, _, _ = tensor_in.size() if d1 == 1: tensor_list = [tensor_upscale(tensor_in,width,height)] else: tensor_list_ = torch.chunk(tensor_in, chunks=d1) tensor_list=[tensor_upscale(i,width,height) for i in tensor_list_] return tensor_list def tensor_upscale(tensor, width, height): samples = tensor.movedim(-1, 1) samples = common_upscale(samples, width, height, "nearest-exact", "center") samples = samples.movedim(1, -1) return samples def nomarl_upscale(img, width, height): samples = img.movedim(-1, 1) img = common_upscale(samples, width, height, "nearest-exact", "center") samples = img.movedim(1, -1) img = tensor2image(samples) return img def cv2tensor(img,bgr2rgb=True): assert type(img) == np.ndarray, 'the img type is {}, but ndarry expected'.format(type(img)) if bgr2rgb: img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img = torch.from_numpy(img.transpose((2, 0, 1))) return img.float().div(255).permute(1, 2, 0).unsqueeze(0) def images_generator(img_list: list, ): # get img size sizes = {} for image_ in img_list: if isinstance(image_, Image.Image): count = sizes.get(image_.size, 0) sizes[image_.size] = count + 1 elif isinstance(image_, np.ndarray): count = sizes.get(image_.shape[:2][::-1], 0) sizes[image_.shape[:2][::-1]] = count + 1 else: raise "unsupport image list,must be pil or cv2!!!" size = max(sizes.items(), key=lambda x: x[1])[0] yield size[0], size[1] # any to tensor def load_image(img_in): if isinstance(img_in, Image.Image): img_in = img_in.convert("RGB") i = np.array(img_in, dtype=np.float32) i = torch.from_numpy(i).div_(255) if i.shape[0] != size[1] or i.shape[1] != size[0]: i = torch.from_numpy(i).movedim(-1, 0).unsqueeze(0) i = common_upscale(i, size[0], size[1], "lanczos", "center") i = i.squeeze(0).movedim(0, -1).numpy() return i elif isinstance(img_in, np.ndarray): i = cv2.cvtColor(img_in, cv2.COLOR_BGR2RGB).astype(np.float32) i = torch.from_numpy(i).div_(255) print(i.shape) return i else: raise "unsupport image list,must be pil,cv2 or tensor!!!" total_images = len(img_list) processed_images = 0 pbar = ProgressBar(total_images) images = map(load_image, img_list) try: prev_image = next(images) while True: next_image = next(images) yield prev_image processed_images += 1 pbar.update_absolute(processed_images, total_images) prev_image = next_image except StopIteration: pass if prev_image is not None: yield prev_image def load_images_list(img_list: list, ): gen = images_generator(img_list) (width, height) = next(gen) images = torch.from_numpy(np.fromiter(gen, np.dtype((np.float32, (height, width, 3))))) if len(images) == 0: raise FileNotFoundError(f"No images could be loaded .") return images def get_video_files(directory, extensions=None): if extensions is None: extensions = ['webm', 'mp4', 'mkv', 'gif', 'mov'] extensions = [ext.lower() for ext in extensions] video_files = [] for root, dirs, files in os.walk(directory): for file in files: _, ext = os.path.splitext(file) ext = ext.lower()[1:] if ext in extensions: full_path = os.path.join(root, file) video_files.append(full_path) return video_files