update readme

This commit is contained in:
chflame163
2024-02-15 16:02:25 +08:00
parent 073c80eef4
commit 48066ea5fa
9 changed files with 32 additions and 14 deletions
+26 -4
View File
@@ -16,17 +16,24 @@ import time
from typing import Union, List
from PIL import Image, ImageFilter, ImageChops, ImageDraw, ImageOps, ImageEnhance, ImageFont
from skimage import img_as_float, img_as_ubyte
from pymatting import fix_trimap, estimate_alpha_cf
from pymatting import fix_trimap, estimate_alpha_cf, estimate_foreground_ml
import torchvision.transforms.functional as TF
import torch.nn.functional as F
import colorsys
from .briarmbg import BriaRMBG
try:
from cv2.ximgproc import guidedFilter
except ImportError:
print(f'# 😺dzNodes: \033[33mDependency package error -> Unable import "guidedFilter", please reinstall "opencv-contrib-python"\033[m')
current_directory = os.path.dirname(os.path.abspath(__file__))
device = "cuda" if torch.cuda.is_available() else "cpu"
def log(message):
def log(message:str, message_type:str='info'):
name = 'LayerStyle'
if message_type == 'error':
message = '\033[33m' + message + '\033[m'
print(f"# 😺dzNodes: {name} -> {message}")
'''Converter'''
@@ -668,8 +675,7 @@ def RMBG(image:Image) -> Image:
return _mask
def mask_edge_detail(image:torch.Tensor, mask:Image, detail_range:int=8, black_point:float=0.01, white_point:float=0.99) -> torch.Tensor:
d = detail_range * 2 + 1
d = detail_range * 5 + 1
i_dup = copy.deepcopy(image.cpu().numpy().astype(np.float64))
a_dup = copy.deepcopy(pil2tensor(mask.convert('RGB')).cpu().numpy().astype(np.float64))
for index, img in enumerate(i_dup):
@@ -682,6 +688,15 @@ def mask_edge_detail(image:torch.Tensor, mask:Image, detail_range:int=8, black_p
a_dup[index] = np.stack([alpha, alpha, alpha], axis=-1) # convert back to rgb
return torch.from_numpy(a_dup.astype(np.float32))
def guided_filter_alpha(image:torch.Tensor, mask:Image, filter_radius:int, sigma:float) -> torch.Tensor:
d = filter_radius + 1
s = sigma / 10
i_dup = copy.deepcopy(image.cpu().numpy())
a_dup = copy.deepcopy(pil2tensor(mask.convert('RGB')).cpu().numpy())
for index, image in enumerate(i_dup):
alpha_work = a_dup[index]
i_dup[index] = guidedFilter(image, alpha_work, d, s)
return torch.from_numpy(i_dup)
def mask_fix(images:torch.Tensor, radius:int, fill_holes:int, white_threshold:float, extra_clip:float) -> torch.Tensor:
d = radius * 2 + 1
@@ -707,6 +722,13 @@ def mask_fix(images:torch.Tensor, radius:int, fill_holes:int, white_threshold:fl
i_dup[index] = cleaned
return torch.from_numpy(i_dup)
def histogram_remap(image:torch.Tensor, blackpoint:float, whitepoint:float) -> torch.Tensor:
bp = min(blackpoint, whitepoint - 0.001)
scale = 1 / (whitepoint - bp)
i_dup = copy.deepcopy(image.cpu().numpy())
i_dup = np.clip((i_dup - bp) * scale, 0.0, 1.0)
return torch.from_numpy(i_dup)
def expand_mask(mask:torch.Tensor, grow:int, blur:int) -> torch.Tensor:
# grow
c = 0