212 lines
6.6 KiB
Python
212 lines
6.6 KiB
Python
# !/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 map_0_1_to_neg1_1(t):
|
||
"""
|
||
接受 torch.Tensor 或可转 torch.Tensor 的输入。
|
||
可处理形状: H,W,C 或 B,H,W,C 或 B,T,H,W,C(会按元素处理)。
|
||
确保 float,0..255 -> 0..1,再把 0..1 -> -1..1(如果已经在 -1..1 则不变)。
|
||
"""
|
||
if not torch.is_tensor(t):
|
||
t = torch.tensor(t)
|
||
t = t.float()
|
||
# 处理 0..255 的情况
|
||
try:
|
||
vmax = float(t.max())
|
||
except Exception:
|
||
vmax = 1.0
|
||
if vmax > 2.0:
|
||
t = t / 255.0
|
||
# 若当前处于 0..1 范围,则映射到 -1..1
|
||
try:
|
||
vmin = float(t.min())
|
||
vmax = float(t.max())
|
||
except Exception:
|
||
vmin, vmax = -1.0, 1.0
|
||
if vmin >= 0.0 and vmax <= 1.1:
|
||
t = t * 2.0 - 1.0
|
||
return t
|
||
|
||
def map_neg1_1_to_0_1(t):
|
||
"""
|
||
接受 torch.Tensor 或可转 torch.Tensor 的输入。
|
||
可处理形状: H,W,C 或 B,H,W,C 或 B,T,H,W,C(会按元素处理)。
|
||
返回 float tensor,范围 0..1。
|
||
"""
|
||
if not torch.is_tensor(t):
|
||
t = torch.tensor(t)
|
||
t = t.float()
|
||
# map -1..1 -> 0..1
|
||
t = (t + 1.0) * 0.5
|
||
# 限幅到 [0,1]
|
||
t = t.clamp(0.0, 1.0)
|
||
# 保持在 cpu 端,调用方可决定是否转 device/dtype
|
||
return t.cpu()
|
||
|
||
|
||
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
|