This commit is contained in:
billwuhao
2025-04-11 10:23:53 +08:00
parent 8792676eac
commit 5df66cd59b
24 changed files with 2962 additions and 2 deletions
+22
View File
@@ -0,0 +1,22 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- master
- main
paths:
- "pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+686
View File
@@ -0,0 +1,686 @@
# ComfyUI-RMBG v2.2.0
# This custom node for ComfyUI provides functionality for background removal using various models,
# including RMBG-2.0, INSPYRENET, BEN, BEN2 and BIREFNET-HR. It leverages deep learning techniques
# to process images and generate masks for background removal.
#
# Models License Notice:
# - RMBG-2.0: Apache-2.0 License (https://huggingface.co/briaai/RMBG-2.0)
# - INSPYRENET: MIT License (https://github.com/plemeri/InSPyReNet)
# - BEN: Apache-2.0 License (https://huggingface.co/PramaLLC/BEN)
# - BEN2: Apache-2.0 License (https://huggingface.co/PramaLLC/BEN2)
#
# This integration script follows GPL-3.0 License.
# When using or modifying this code, please respect both the original model licenses
# and this integration's license terms.
#
# Source: https://github.com/1038lab/ComfyUI-RMBG
import os
import torch
from PIL import Image
from torchvision import transforms
import numpy as np
import folder_paths
from PIL import ImageFilter
import torch.nn.functional as F
from huggingface_hub import hf_hub_download
import shutil
import sys
import importlib.util
from transformers import AutoModelForImageSegmentation
import cv2
import types
device = "cuda" if torch.cuda.is_available() else "cpu"
# Add model path
folder_paths.add_model_folder_path("rmbg", os.path.join(folder_paths.models_dir, "RMBG"))
# Model configuration
AVAILABLE_MODELS = {
"RMBG-2.0": {
"type": "rmbg",
"repo_id": "1038lab/RMBG-2.0",
"files": {
"config.json": "config.json",
"model.safetensors": "model.safetensors",
"birefnet.py": "birefnet.py",
"BiRefNet_config.py": "BiRefNet_config.py"
},
"cache_dir": "RMBG-2.0"
},
"INSPYRENET": {
"type": "inspyrenet",
"repo_id": "1038lab/inspyrenet",
"files": {
"inspyrenet.safetensors": "inspyrenet.safetensors"
},
"cache_dir": "INSPYRENET"
},
"BEN": {
"type": "ben",
"repo_id": "1038lab/BEN",
"files": {
"model.py": "model.py",
"BEN_Base.pth": "BEN_Base.pth"
},
"cache_dir": "BEN"
},
"BEN2": {
"type": "ben2",
"repo_id": "1038lab/BEN2",
"files": {
"BEN2_Base.pth": "BEN2_Base.pth",
"BEN2.py": "BEN2.py"
},
"cache_dir": "BEN2"
}
}
# Utility functions
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
def pil2tensor(image):
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
def handle_model_error(message):
print(f"[RMBG ERROR] {message}")
raise RuntimeError(message)
class BaseModelLoader:
def __init__(self):
self.model = None
self.current_model_version = None
self.base_cache_dir = os.path.join(folder_paths.models_dir, "RMBG")
def get_cache_dir(self, model_name):
cache_path = os.path.join(self.base_cache_dir, AVAILABLE_MODELS[model_name]["cache_dir"])
os.makedirs(cache_path, exist_ok=True)
return cache_path
def check_model_cache(self, model_name):
model_info = AVAILABLE_MODELS[model_name]
cache_dir = self.get_cache_dir(model_name)
if not os.path.exists(cache_dir):
return False, "Model directory not found"
missing_files = []
for filename in model_info["files"].keys():
if not os.path.exists(os.path.join(cache_dir, model_info["files"][filename])):
missing_files.append(filename)
if missing_files:
return False, f"Missing model files: {', '.join(missing_files)}"
return True, "Model cache verified"
def download_model(self, model_name):
model_info = AVAILABLE_MODELS[model_name]
cache_dir = self.get_cache_dir(model_name)
try:
os.makedirs(cache_dir, exist_ok=True)
print(f"Downloading {model_name} model files...")
for filename in model_info["files"].keys():
print(f"Downloading {filename}...")
hf_hub_download(
repo_id=model_info["repo_id"],
filename=filename,
local_dir=cache_dir
)
return True, "Model files downloaded successfully"
except Exception as e:
return False, f"Error downloading model files: {str(e)}"
def clear_model(self):
if self.model is not None:
self.model.cpu()
del self.model
import gc
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
self.model = None
self.current_model_version = None
class RMBGModel(BaseModelLoader):
def __init__(self):
super().__init__()
def load_model(self, model_name):
if self.current_model_version != model_name:
self.clear_model()
cache_dir = self.get_cache_dir(model_name)
try:
# Try standard loading first
try:
self.model = AutoModelForImageSegmentation.from_pretrained(
cache_dir,
trust_remote_code=True,
local_files_only=True
)
except AttributeError as ae:
if "'Config' object has no attribute 'get_text_config'" in str(ae):
print("[RMBG WARNING] Detected newer transformers version, using compatibility mode...")
try:
from transformers import PreTrainedModel
import json
config_path = os.path.join(cache_dir, "config.json")
with open(config_path, 'r') as f:
config = json.load(f)
birefnet_path = os.path.join(cache_dir, "birefnet.py")
BiRefNetConfig_path = os.path.join(cache_dir, "BiRefNet_config.py")
# Load the BiRefNetConfig
config_spec = importlib.util.spec_from_file_location("BiRefNetConfig", BiRefNetConfig_path)
config_module = importlib.util.module_from_spec(config_spec)
sys.modules["BiRefNetConfig"] = config_module
config_spec.loader.exec_module(config_module)
# Fix and load birefnet module
with open(birefnet_path, 'r') as f:
birefnet_content = f.read()
birefnet_content = birefnet_content.replace(
"from .BiRefNet_config import BiRefNetConfig",
"from BiRefNetConfig import BiRefNetConfig"
)
module_name = f"custom_birefnet_model_{hash(birefnet_path)}"
module = types.ModuleType(module_name)
sys.modules[module_name] = module
exec(birefnet_content, module.__dict__)
for attr_name in dir(module):
attr = getattr(module, attr_name)
if isinstance(attr, type) and issubclass(attr, PreTrainedModel) and attr != PreTrainedModel:
BiRefNetConfig = getattr(config_module, "BiRefNetConfig")
model_config = BiRefNetConfig()
self.model = attr(model_config)
weights_path = os.path.join(cache_dir, "model.safetensors")
try:
try:
import safetensors.torch
self.model.load_state_dict(safetensors.torch.load_file(weights_path))
except ImportError:
from transformers.modeling_utils import load_state_dict
state_dict = load_state_dict(weights_path)
self.model.load_state_dict(state_dict)
except Exception as load_error:
pytorch_weights = os.path.join(cache_dir, "pytorch_model.bin")
if os.path.exists(pytorch_weights):
self.model.load_state_dict(torch.load(pytorch_weights, map_location="cpu"))
else:
raise RuntimeError(f"Failed to load weights: {str(load_error)}")
break
if self.model is None:
raise RuntimeError("Could not find suitable model class")
except Exception as custom_e:
handle_model_error(f"Failed to load model in compatibility mode: {str(custom_e)}")
else:
raise ae
except Exception as e:
handle_model_error(f"Error loading model: {str(e)}")
self.model.eval()
for param in self.model.parameters():
param.requires_grad = False
torch.set_float32_matmul_precision('high')
self.model.to(device)
self.current_model_version = model_name
def process_image(self, images, model_name, params):
try:
self.load_model(model_name)
# Prepare batch processing
transform_image = transforms.Compose([
transforms.Resize((params["process_res"], params["process_res"])),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
# Ensure input is in list format
if isinstance(images, torch.Tensor):
if len(images.shape) == 3:
images = [images]
else:
images = [img for img in images]
# Store original image sizes
original_sizes = [tensor2pil(img).size for img in images]
# Batch process transformations
input_tensors = [transform_image(tensor2pil(img)).unsqueeze(0) for img in images]
input_batch = torch.cat(input_tensors, dim=0).to(device)
with torch.no_grad():
outputs = self.model(input_batch)
if isinstance(outputs, list) and len(outputs) > 0:
results = outputs[-1].sigmoid().cpu()
elif isinstance(outputs, dict) and 'logits' in outputs:
results = outputs['logits'].sigmoid().cpu()
elif isinstance(outputs, torch.Tensor):
results = outputs.sigmoid().cpu()
else:
try:
if hasattr(outputs, 'last_hidden_state'):
results = outputs.last_hidden_state.sigmoid().cpu()
else:
for k, v in outputs.items():
if isinstance(v, torch.Tensor):
results = v.sigmoid().cpu()
break
except:
handle_model_error("Unable to recognize model output format")
masks = []
# Process each result and resize back to original dimensions
for i, (result, (orig_w, orig_h)) in enumerate(zip(results, original_sizes)):
result = result.squeeze()
result = result * (1 + (1 - params["sensitivity"]))
result = torch.clamp(result, 0, 1)
# Resize back to original dimensions
result = F.interpolate(result.unsqueeze(0).unsqueeze(0),
size=(orig_h, orig_w),
mode='bilinear').squeeze()
masks.append(tensor2pil(result))
return masks
except Exception as e:
handle_model_error(f"Error in batch processing: {str(e)}")
class InspyrenetModel(BaseModelLoader):
def __init__(self):
super().__init__()
def load_model(self, model_name):
if self.current_model_version != model_name:
self.clear_model()
try:
import transparent_background
self.model = transparent_background.Remover()
self.current_model_version = model_name
except ImportError:
try:
import pip
pip.main(['install', 'transparent_background'])
import transparent_background
self.model = transparent_background.Remover()
self.current_model_version = model_name
except Exception as e:
handle_model_error(f"Failed to install transparent_background: {str(e)}")
def process_image(self, image, model_name, params):
try:
self.load_model(model_name)
orig_image = tensor2pil(image)
w, h = orig_image.size
# Resize for processing
aspect_ratio = h / w
new_w = params["process_res"]
new_h = int(params["process_res"] * aspect_ratio)
resized_image = orig_image.resize((new_w, new_h), Image.LANCZOS)
# Process image
foreground = self.model.process(resized_image, type='rgba')
foreground = foreground.resize((w, h), Image.LANCZOS)
mask = foreground.split()[-1]
return mask
except Exception as e:
handle_model_error(f"Error in Inspyrenet processing: {str(e)}")
class BENModel(BaseModelLoader):
def __init__(self):
super().__init__()
def load_model(self, model_name):
if self.current_model_version != model_name:
self.clear_model()
cache_dir = self.get_cache_dir(model_name)
model_path = os.path.join(cache_dir, "model.py")
module_name = f"custom_ben_model_{hash(model_path)}"
spec = importlib.util.spec_from_file_location(module_name, model_path)
ben_module = importlib.util.module_from_spec(spec)
sys.modules[module_name] = ben_module
spec.loader.exec_module(ben_module)
model_weights_path = os.path.join(cache_dir, "BEN_Base.pth")
self.model = ben_module.BEN_Base()
self.model.loadcheckpoints(model_weights_path)
self.model.eval()
for param in self.model.parameters():
param.requires_grad = False
torch.set_float32_matmul_precision('high')
self.model.to(device)
self.current_model_version = model_name
def process_image(self, image, model_name, params):
try:
self.load_model(model_name)
orig_image = tensor2pil(image)
w, h = orig_image.size
aspect_ratio = h / w
new_w = params["process_res"]
new_h = int(params["process_res"] * aspect_ratio)
resized_image = orig_image.resize((new_w, new_h), Image.LANCZOS)
processed_input = resized_image.convert("RGBA")
with torch.no_grad():
_, foreground = self.model.inference(processed_input)
foreground = foreground.resize((w, h), Image.LANCZOS)
mask = foreground.split()[-1]
return mask
except Exception as e:
handle_model_error(f"Error in BEN processing: {str(e)}")
class BEN2Model(BaseModelLoader):
def __init__(self):
super().__init__()
def load_model(self, model_name):
if self.current_model_version != model_name:
self.clear_model()
try:
cache_dir = self.get_cache_dir(model_name)
model_path = os.path.join(cache_dir, "BEN2.py")
module_name = f"custom_ben2_model_{hash(model_path)}"
spec = importlib.util.spec_from_file_location(module_name, model_path)
ben2_module = importlib.util.module_from_spec(spec)
sys.modules[module_name] = ben2_module
spec.loader.exec_module(ben2_module)
model_weights_path = os.path.join(cache_dir, "BEN2_Base.pth")
self.model = ben2_module.BEN_Base()
self.model.loadcheckpoints(model_weights_path)
self.model.eval()
for param in self.model.parameters():
param.requires_grad = False
torch.set_float32_matmul_precision('high')
self.model.to(device)
self.current_model_version = model_name
except Exception as e:
handle_model_error(f"Error loading BEN2 model: {str(e)}")
def process_image(self, images, model_name, params):
try:
self.load_model(model_name)
if isinstance(images, torch.Tensor):
if len(images.shape) == 3:
images = [images]
else:
images = [img for img in images]
batch_size = 3
all_masks = []
for i in range(0, len(images), batch_size):
batch_images = images[i:i + batch_size]
batch_pil_images = []
original_sizes = []
for img in batch_images:
orig_image = tensor2pil(img)
w, h = orig_image.size
original_sizes.append((w, h))
aspect_ratio = h / w
new_w = params["process_res"]
new_h = int(params["process_res"] * aspect_ratio)
resized_image = orig_image.resize((new_w, new_h), Image.LANCZOS)
processed_input = resized_image.convert("RGBA")
batch_pil_images.append(processed_input)
with torch.no_grad():
try:
foregrounds = self.model.inference(batch_pil_images)
if not isinstance(foregrounds, list):
foregrounds = [foregrounds]
except Exception as e:
handle_model_error(f"Error in BEN2 inference: {str(e)}")
for foreground, (orig_w, orig_h) in zip(foregrounds, original_sizes):
foreground = foreground.resize((orig_w, orig_h), Image.LANCZOS)
mask = foreground.split()[-1]
all_masks.append(mask)
if len(all_masks) == 1:
return all_masks[0]
return all_masks
except Exception as e:
handle_model_error(f"Error in BEN2 processing: {str(e)}")
def refine_foreground(image_bchw, masks_b1hw):
b, c, h, w = image_bchw.shape
if b != masks_b1hw.shape[0]:
raise ValueError("images and masks must have the same batch size")
image_np = image_bchw.cpu().numpy()
mask_np = masks_b1hw.cpu().numpy()
refined_fg = []
for i in range(b):
mask = mask_np[i, 0]
thresh = 0.45
mask_binary = (mask > thresh).astype(np.float32)
edge_blur = cv2.GaussianBlur(mask_binary, (3, 3), 0)
transition_mask = np.logical_and(mask > 0.05, mask < 0.95)
alpha = 0.85
mask_refined = np.where(transition_mask,
alpha * mask + (1-alpha) * edge_blur,
mask_binary)
edge_region = np.logical_and(mask > 0.2, mask < 0.8)
mask_refined = np.where(edge_region,
mask_refined * 0.98,
mask_refined)
result = []
for c in range(image_np.shape[1]):
channel = image_np[i, c]
refined = channel * mask_refined
result.append(refined)
refined_fg.append(np.stack(result))
return torch.from_numpy(np.stack(refined_fg))
class RMBG:
def __init__(self):
self.models = {
"RMBG-2.0": RMBGModel(),
"INSPYRENET": InspyrenetModel(),
"BEN": BENModel(),
"BEN2": BEN2Model()
}
@classmethod
def INPUT_TYPES(s):
tooltips = {
"image": "Input image to be processed for background removal.",
"model": "Select the background removal model to use (RMBG-2.0, INSPYRENET, BEN).",
"sensitivity": "Adjust the strength of mask detection (higher values result in more aggressive detection).",
"process_res": "Set the processing resolution (higher values require more VRAM and may increase processing time).",
"mask_blur": "Specify the amount of blur to apply to the mask edges (0 for no blur, higher values for more blur).",
"mask_offset": "Adjust the mask boundary (positive values expand the mask, negative values shrink it).",
"background": "Choose the background color for the final output (Alpha for transparent background).",
"invert_output": "Enable to invert both the image and mask output (useful for certain effects).",
"optimize": "Enable model optimization for faster processing (may affect output quality).",
"refine_foreground": "Use Fast Foreground Colour Estimation to optimize transparent background"
}
return {
"required": {
"image": ("IMAGE", {"tooltip": tooltips["image"]}),
"model": (list(AVAILABLE_MODELS.keys()), {"tooltip": tooltips["model"]}),
},
"optional": {
"sensitivity": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": tooltips["sensitivity"]}),
"process_res": ("INT", {"default": 1024, "min": 256, "max": 2048, "step": 8, "tooltip": tooltips["process_res"]}),
"mask_blur": ("INT", {"default": 0, "min": 0, "max": 64, "step": 1, "tooltip": tooltips["mask_blur"]}),
"mask_offset": ("INT", {"default": 0, "min": -64, "max": 64, "step": 1, "tooltip": tooltips["mask_offset"]}),
"background": (["Alpha", "black", "white", "gray", "green", "blue", "red"], {"default": "Alpha", "tooltip": tooltips["background"]}),
"invert_output": ("BOOLEAN", {"default": False, "tooltip": tooltips["invert_output"]}),
"optimize": (["default", "on"], {"default": "default", "tooltip": tooltips["optimize"]}),
"refine_foreground": ("BOOLEAN", {"default": False, "tooltip": tooltips["refine_foreground"]})
}
}
RETURN_TYPES = ("IMAGE", "MASK", "IMAGE")
RETURN_NAMES = ("IMAGE", "MASK", "MASK_IMAGE")
FUNCTION = "process_image"
CATEGORY = "🧪AILab/🧽RMBG"
def process_image(self, image, model, **params):
try:
processed_images = []
processed_masks = []
bg_colors = {
"Alpha": None,
"black": (0, 0, 0),
"white": (255, 255, 255),
"gray": (128, 128, 128),
"green": (0, 255, 0),
"blue": (0, 0, 255),
"red": (255, 0, 0)
}
model_instance = self.models[model]
# Check and download model if needed
cache_status, message = model_instance.check_model_cache(model)
if not cache_status:
print(f"Cache check: {message}")
print("Downloading required model files...")
download_status, download_message = model_instance.download_model(model)
if not download_status:
handle_model_error(download_message)
print("Model files downloaded successfully")
for img in image:
# Get mask from specific model
mask = model_instance.process_image(img, model, params)
# Ensure mask is in the correct format
if isinstance(mask, list):
masks = [m.convert("L") for m in mask if isinstance(m, Image.Image)]
mask = masks[0] if masks else None
elif isinstance(mask, Image.Image):
mask = mask.convert("L")
# Post-process mask
mask_tensor = pil2tensor(mask)
mask_tensor = mask_tensor * (1 + (1 - params["sensitivity"]))
mask_tensor = torch.clamp(mask_tensor, 0, 1)
mask = tensor2pil(mask_tensor)
if params["mask_blur"] > 0:
mask = mask.filter(ImageFilter.GaussianBlur(radius=params["mask_blur"]))
if params["mask_offset"] != 0:
if params["mask_offset"] > 0:
for _ in range(params["mask_offset"]):
mask = mask.filter(ImageFilter.MaxFilter(3))
else:
for _ in range(-params["mask_offset"]):
mask = mask.filter(ImageFilter.MinFilter(3))
if params["invert_output"]:
mask = Image.fromarray(255 - np.array(mask))
# Convert to tensors for refine_foreground
img_tensor = torch.from_numpy(np.array(tensor2pil(img))).permute(2, 0, 1).unsqueeze(0) / 255.0
mask_tensor = torch.from_numpy(np.array(mask)).unsqueeze(0).unsqueeze(0) / 255.0
# Create final image
orig_image = tensor2pil(img)
if params.get("refine_foreground", False):
refined_fg = refine_foreground(img_tensor, mask_tensor)
refined_fg = tensor2pil(refined_fg[0].permute(1, 2, 0))
r, g, b = refined_fg.split()
foreground = Image.merge('RGBA', (r, g, b, mask))
else:
orig_rgba = orig_image.convert("RGBA")
r, g, b, _ = orig_rgba.split()
foreground = Image.merge('RGBA', (r, g, b, mask))
if params["background"] != "Alpha":
bg_color = bg_colors[params["background"]]
bg_image = Image.new('RGBA', orig_image.size, (*bg_color, 255))
composite_image = Image.alpha_composite(bg_image, foreground)
processed_images.append(pil2tensor(composite_image.convert("RGB")))
else:
processed_images.append(pil2tensor(foreground))
processed_masks.append(pil2tensor(mask))
# Create mask image for visualization
mask_images = []
for mask_tensor in processed_masks:
# Convert mask to RGB image format for visualization
mask_image = mask_tensor.reshape((-1, 1, mask_tensor.shape[-2], mask_tensor.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3)
mask_images.append(mask_image)
mask_image_output = torch.cat(mask_images, dim=0)
return (torch.cat(processed_images, dim=0), torch.cat(processed_masks, dim=0), mask_image_output)
except Exception as e:
handle_model_error(f"Error in image processing: {str(e)}")
# Return original image and empty mask on error
empty_mask = torch.zeros((image.shape[0], image.shape[2], image.shape[3]))
empty_mask_image = empty_mask.reshape((-1, 1, empty_mask.shape[-2], empty_mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3)
return (image, empty_mask, empty_mask_image)
# Node Mapping
NODE_CLASS_MAPPINGS = {
"RMBG": RMBG
}
NODE_DISPLAY_NAME_MAPPINGS = {
"RMBG": "Remove Background (RMBG)"
}
+55
View File
@@ -0,0 +1,55 @@
[中文](README-CN.md)|[English](README.md)
# 人像处理等相关的 ComfyUI 节点
目前包含以下节点:
- 图片人脸对齐(正脸);
- 人脸检测裁剪, 可选是否对齐, 可调裁剪区域大小, 角度;
- 各种证件照一键生成;
- 美化照片, 包括亮度, 饱和度, 锐化, 磨皮等.
示例:
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_07-06-36.png)
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_07-08-46.png)
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_09-05-41.png)
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_09-27-16.png)
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_09-48-23.png)
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_07-10-24.png)
## 📣 更新
[2025-04-11]⚒️: 发布版本 v1.0.0.
## 安装
```
cd ComfyUI/custom_nodes
git clone https://github.com/billwuhao/ComfyUI_PortraitTools.git
cd ComfyUI_PortraitTools
pip install -r requirements.txt
# python_embeded
./python_embeded/python.exe -m pip install -r requirements.txt
```
## 模型下载
如果你正在使用 [ComfyUI-ReActor](https://github.com/Gourieff/comfyui-reactor) 和 [ComfyUI-RMBG](https://github.com/1038lab/ComfyUI-RMBG) 节点, 不用下载模型, 它们是公用的.
否则, 下载 [detection_Resnet50_Final.pth](https://huggingface.co/salmonrk/facedetection/blob/main/detection_Resnet50_Final.pth) 放到 `ComfyUI\models\facedetection` 文件夹下. [ComfyUI-RMBG](https://github.com/1038lab/ComfyUI-RMBG) 的模型会自动下载到 `ComfyUI\models\RMBG` 文件夹下.
## 鸣谢
感谢以下项目:
- [HivisionIDPhotos](https://github.com/Zeyi-Lin/HivisionIDPhotos)
- [ComfyUI-RMBG](https://github.com/1038lab/ComfyUI-RMBG)
- [HivisionIDPhotos-ComfyUI](https://github.com/AIFSH/HivisionIDPhotos-ComfyUI)
- [facerestore_cf](https://github.com/mav-rik/facerestore_cf)
+55 -2
View File
@@ -1,2 +1,55 @@
# ComfyUI_PortraitTools
Portrait Tools: Facial detection cropping, alignment, ID photo, etc
[中文](README-CN.md) | [English](README.md)
# ComfyUI Nodes for Portrait Processing
Currently includes the following nodes:
- Image face alignment (frontal);
- Face detection and cropping, with optional alignment, adjustable crop area size, and angle;
- One-click generation of various passport photos;
- Photo enhancement, including brightness, saturation, sharpening, and skin smoothing.
Examples:
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_07-06-36.png)
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_07-08-46.png)
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_09-05-41.png)
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_09-27-16.png)
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_09-48-23.png)
![](https://github.com/billwuhao/ComfyUI_PortraitTools/blob/main/images/2025-04-11_07-10-24.png)
## 📣 Updates
[2025-04-11] ⚒️: Released version v1.0.0.
## Installation
```
cd ComfyUI/custom_nodes
git clone https://github.com/billwuhao/ComfyUI_PortraitTools.git
cd ComfyUI_PortraitTools
pip install -r requirements.txt
# python_embeded
./python_embeded/python.exe -m pip install -r requirements.txt
```
## Model Download
If you are using the [ComfyUI-ReActor](https://github.com/Gourieff/comfyui-reactor) and [ComfyUI-RMBG](https://github.com/1038lab/ComfyUI-RMBG) nodes, you do not need to download the models, as they are shared.
Otherwise, download [detection_Resnet50_Final.pth](https://huggingface.co/salmonrk/facedetection/blob/main/detection_Resnet50_Final.pth) and place it in the `ComfyUI\models\facedetection` folder. The models for [ComfyUI-RMBG](https://github.com/1038lab/ComfyUI-RMBG) will be automatically downloaded to the `ComfyUI\models\RMBG` folder.
## Acknowledgements
Thanks to the following projects:
- [HivisionIDPhotos](https://github.com/Zeyi-Lin/HivisionIDPhotos)
- [ComfyUI-RMBG](https://github.com/1038lab/ComfyUI-RMBG)
- [HivisionIDPhotos-ComfyUI](https://github.com/AIFSH/HivisionIDPhotos-ComfyUI)
- [facerestore_cf](https://github.com/mav-rik/facerestore_cf)
+3
View File
@@ -0,0 +1,3 @@
from .ptnodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+3
View File
@@ -0,0 +1,3 @@
from beauty.base_adjust import *
from beauty.grind_skin import *
from beauty.whitening import *
+103
View File
@@ -0,0 +1,103 @@
"""
亮度、对比度、锐化、饱和度调整模块
"""
import cv2
import numpy as np
def adjust_brightness_contrast_sharpen_saturation(
image,
brightness_factor=0,
contrast_factor=0,
sharpen_strength=0,
saturation_factor=0,
):
"""
调整图像的亮度、对比度、锐度和饱和度。
参数:
image (numpy.ndarray): 输入的图像数组。
brightness_factor (float): 亮度调整因子。大于0增加亮度,小于0降低亮度。
contrast_factor (float): 对比度调整因子。大于0增加对比度,小于0降低对比度。
sharpen_strength (float): 锐化强度。
saturation_factor (float): 饱和度调整因子。大于0增加饱和度,小于0降低饱和度。
返回:
numpy.ndarray: 调整后的图像。
"""
if (
brightness_factor == 0
and contrast_factor == 0
and sharpen_strength == 0
and saturation_factor == 0
):
return image.copy()
adjusted_image = image.copy()
# 调整饱和度
if saturation_factor != 0:
adjusted_image = adjust_saturation(adjusted_image, saturation_factor)
# 调整亮度和对比度
alpha = 1.0 + (contrast_factor / 100.0)
beta = brightness_factor
adjusted_image = cv2.convertScaleAbs(adjusted_image, alpha=alpha, beta=beta)
# 增强锐化
adjusted_image = sharpen_image(adjusted_image, sharpen_strength)
return adjusted_image
def adjust_saturation(image, saturation_factor):
"""
调整图像的饱和度。
参数:
image (numpy.ndarray): 输入的图像数组。
saturation_factor (float): 饱和度调整因子。大于0增加饱和度,小于0降低饱和度。
返回:
numpy.ndarray: 调整后的图像。
"""
hsv = cv2.cvtColor(image, cv2.COLOR_BGR2HSV)
h, s, v = cv2.split(hsv)
s = s.astype(np.float32)
s = s + s * (saturation_factor / 100.0)
s = np.clip(s, 0, 255).astype(np.uint8)
hsv = cv2.merge([h, s, v])
return cv2.cvtColor(hsv, cv2.COLOR_HSV2BGR)
def sharpen_image(image, strength=0):
"""
对图像进行锐化处理。
参数:
image (numpy.ndarray): 输入的图像数组。
strength (float): 锐化强度,范围建议为0-5。0表示不进行锐化。
返回:
numpy.ndarray: 锐化后的图像。
"""
print(f"Sharpen strength: {strength}")
if strength == 0:
return image.copy()
strength = strength * 20
kernel_strength = 1 + (strength / 500)
kernel = (
np.array([[-0.5, -0.5, -0.5], [-0.5, 5, -0.5], [-0.5, -0.5, -0.5]])
* kernel_strength
)
sharpened = cv2.filter2D(image, -1, kernel)
sharpened = np.clip(sharpened, 0, 255).astype(np.uint8)
alpha = strength / 200
blended = cv2.addWeighted(image, 1 - alpha, sharpened, alpha, 0)
return blended
+30
View File
@@ -0,0 +1,30 @@
# Required Libraries
import cv2
def grindSkin(src, grindDegree: int = 3, detailDegree: int = 1, strength: int = 9):
"""
Dest =(Src * (100 - Opacity) + (Src + 2 * GaussBlur(EPFFilter(Src) - Src)) * Opacity) / 100
人像磨皮方案
Args:
src: 原图
grindDegree: 磨皮程度调节参数
detailDegree: 细节程度调节参数
strength: 融合程度,作为磨皮强度(0 - 10)
Returns:
磨皮后的图像
"""
if strength <= 0:
return src
dst = src.copy()
opacity = min(10.0, strength) / 10.0
dx = grindDegree * 5
fc = grindDegree * 12.5
temp1 = cv2.bilateralFilter(src[:, :, :3], dx, fc, fc)
temp2 = cv2.subtract(temp1, src[:, :, :3])
temp3 = cv2.GaussianBlur(temp2, (2 * detailDegree - 1, 2 * detailDegree - 1), 0)
temp4 = cv2.add(cv2.add(temp3, temp3), src[:, :, :3])
dst[:, :, :3] = cv2.addWeighted(temp4, opacity, src[:, :, :3], 1 - opacity, 0.0)
return dst
Binary file not shown.

After

Width:  |  Height:  |  Size: 103 KiB

+75
View File
@@ -0,0 +1,75 @@
import cv2
import numpy as np
import os
class LutWhite:
CUBE64_ROWS = 8
CUBE64_SIZE = 64
CUBE256_SIZE = 256
CUBE_SCALE = CUBE256_SIZE // CUBE64_SIZE
def __init__(self, lut_image):
self.lut = self._create_lut(lut_image)
def _create_lut(self, lut_image):
reshape_lut = np.zeros(
(self.CUBE256_SIZE, self.CUBE256_SIZE, self.CUBE256_SIZE, 3), dtype=np.uint8
)
for i in range(self.CUBE64_SIZE):
tmp = i // self.CUBE64_ROWS
cx = (i % self.CUBE64_ROWS) * self.CUBE64_SIZE
cy = tmp * self.CUBE64_SIZE
cube64 = lut_image[cy : cy + self.CUBE64_SIZE, cx : cx + self.CUBE64_SIZE]
if cube64.size == 0:
continue
cube256 = cv2.resize(cube64, (self.CUBE256_SIZE, self.CUBE256_SIZE))
reshape_lut[i * self.CUBE_SCALE : (i + 1) * self.CUBE_SCALE] = cube256
return reshape_lut
def apply(self, src):
b, g, r = src[:, :, 0], src[:, :, 1], src[:, :, 2]
return self.lut[b, g, r]
class MakeWhiter:
def __init__(self, lut_image):
self.lut_white = LutWhite(lut_image)
def run(self, src: np.ndarray, strength: int) -> np.ndarray:
strength = np.clip(strength / 10.0, 0, 1)
if strength <= 0:
return src
img = self.lut_white.apply(src[:, :, :3])
return cv2.addWeighted(src[:, :, :3], 1 - strength, img, strength, 0)
base_dir = os.path.dirname(os.path.abspath(__file__))
default_lut = cv2.imread(os.path.join(base_dir, "lut/lut_origin.png"))
make_whiter = MakeWhiter(default_lut)
def make_whitening(image, strength):
image = cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
iteration = strength // 10
bias = strength % 10
for i in range(iteration):
image = make_whiter.run(image, 10)
image = make_whiter.run(image, bias)
return cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
def make_whitening_png(image, strength):
image = cv2.cvtColor(np.array(image), cv2.COLOR_RGBA2BGRA)
b, g, r, a = cv2.split(image)
bgr_image = cv2.merge((b, g, r))
b_w, g_w, r_w = cv2.split(make_whiter.run(bgr_image, strength))
output_image = cv2.merge((b_w, g_w, r_w, a))
return cv2.cvtColor(output_image, cv2.COLOR_RGBA2BGRA)
Binary file not shown.

After

Width:  |  Height:  |  Size: 174 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 270 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 190 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 406 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 222 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 398 KiB

+140
View File
@@ -0,0 +1,140 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
r"""
@DATE: 2024/9/5 21:35
@File: layout_calculator.py
@IDE: pycharm
@Description:
布局计算器
"""
import cv2
import numpy as np
def judge_layout(
input_width,
input_height,
PHOTO_INTERVAL_W,
PHOTO_INTERVAL_H,
LIMIT_BLOCK_W,
LIMIT_BLOCK_H,
):
centerBlockHeight_1, centerBlockWidth_1 = (
input_height,
input_width,
) # 由证件照们组成的一个中心区块(1 代表不转置排列)
centerBlockHeight_2, centerBlockWidth_2 = (
input_width,
input_height,
) # 由证件照们组成的一个中心区块(2 代表转置排列)
# 1.不转置排列的情况下:
layout_col_no_transpose = 0 # 行
layout_row_no_transpose = 0 # 列
for i in range(1, 4):
centerBlockHeight_temp = input_height * i + PHOTO_INTERVAL_H * (i - 1)
if centerBlockHeight_temp < LIMIT_BLOCK_H:
centerBlockHeight_1 = centerBlockHeight_temp
layout_row_no_transpose = i
else:
break
for j in range(1, 9):
centerBlockWidth_temp = input_width * j + PHOTO_INTERVAL_W * (j - 1)
if centerBlockWidth_temp < LIMIT_BLOCK_W:
centerBlockWidth_1 = centerBlockWidth_temp
layout_col_no_transpose = j
else:
break
layout_number_no_transpose = layout_row_no_transpose * layout_col_no_transpose
# 2.转置排列的情况下:
layout_col_transpose = 0 # 行
layout_row_transpose = 0 # 列
for i in range(1, 4):
centerBlockHeight_temp = input_width * i + PHOTO_INTERVAL_H * (i - 1)
if centerBlockHeight_temp < LIMIT_BLOCK_H:
centerBlockHeight_2 = centerBlockHeight_temp
layout_row_transpose = i
else:
break
for j in range(1, 9):
centerBlockWidth_temp = input_height * j + PHOTO_INTERVAL_W * (j - 1)
if centerBlockWidth_temp < LIMIT_BLOCK_W:
centerBlockWidth_2 = centerBlockWidth_temp
layout_col_transpose = j
else:
break
layout_number_transpose = layout_row_transpose * layout_col_transpose
if layout_number_transpose > layout_number_no_transpose:
layout_mode = (layout_col_transpose, layout_row_transpose, 2)
return layout_mode, centerBlockWidth_2, centerBlockHeight_2
else:
layout_mode = (layout_col_no_transpose, layout_row_no_transpose, 1)
return layout_mode, centerBlockWidth_1, centerBlockHeight_1
def generate_layout_photo(input_height, input_width):
# 1.基础参数表
LAYOUT_WIDTH = 1746
LAYOUT_HEIGHT = 1180
PHOTO_INTERVAL_H = 30 # 证件照与证件照之间的垂直距离
PHOTO_INTERVAL_W = 30 # 证件照与证件照之间的水平距离
SIDES_INTERVAL_H = 50 # 证件照与画布边缘的垂直距离
SIDES_INTERVAL_W = 70 # 证件照与画布边缘的水平距离
LIMIT_BLOCK_W = LAYOUT_WIDTH - 2 * SIDES_INTERVAL_W
LIMIT_BLOCK_H = LAYOUT_HEIGHT - 2 * SIDES_INTERVAL_H
# 2.创建一个 1180x1746 的空白画布
white_background = np.zeros([LAYOUT_HEIGHT, LAYOUT_WIDTH, 3], np.uint8)
white_background.fill(255)
# 3.计算照片的 layout(列、行、横竖朝向),证件照组成的中心区块的分辨率
layout_mode, centerBlockWidth, centerBlockHeight = judge_layout(
input_width,
input_height,
PHOTO_INTERVAL_W,
PHOTO_INTERVAL_H,
LIMIT_BLOCK_W,
LIMIT_BLOCK_H,
)
# 4.开始排列组合
x11 = (LAYOUT_WIDTH - centerBlockWidth) // 2
y11 = (LAYOUT_HEIGHT - centerBlockHeight) // 2
typography_arr = []
typography_rotate = False
if layout_mode[2] == 2:
input_height, input_width = input_width, input_height
typography_rotate = True
for j in range(layout_mode[1]):
for i in range(layout_mode[0]):
xi = x11 + i * input_width + i * PHOTO_INTERVAL_W
yi = y11 + j * input_height + j * PHOTO_INTERVAL_H
typography_arr.append([xi, yi])
return typography_arr, typography_rotate
def generate_layout_image(
input_image, typography_arr, typography_rotate, width=295, height=413
):
LAYOUT_WIDTH = 1746
LAYOUT_HEIGHT = 1180
white_background = np.zeros([LAYOUT_HEIGHT, LAYOUT_WIDTH, 3], np.uint8)
white_background.fill(255)
if input_image.shape[0] != height:
input_image = cv2.resize(input_image, (width, height))
if typography_rotate:
input_image = cv2.transpose(input_image)
input_image = cv2.flip(input_image, 0) # 0 表示垂直镜像
height, width = width, height
for arr in typography_arr:
locate_x, locate_y = arr[0], arr[1]
white_background[locate_y : locate_y + height, locate_x : locate_x + width] = (
input_image
)
return white_background
+676
View File
@@ -0,0 +1,676 @@
import os
import torch
from copy import deepcopy
import cv2
import numpy as np
from comfy import model_management
import folder_paths
import sys
from PIL import Image
current_dir = os.path.dirname(os.path.abspath(__file__))
if current_dir not in sys.path:
sys.path.append(current_dir)
from retinaface import RetinaFace
from layout_calculator import generate_layout_photo, generate_layout_image
from beauty import grindSkin, make_whitening, adjust_brightness_contrast_sharpen_saturation
from AILab_RMBG import (AVAILABLE_MODELS,
RMBGModel,
BENModel,
BEN2Model,
InspyrenetModel,
tensor2pil,
pil2tensor,
handle_model_error,
)
models_dir = folder_paths.models_dir
model_path = os.path.join(models_dir, "facedetection", "detection_Resnet50_Final.pth")
device = model_management.get_torch_device()
def init_model(half=False, device=device):
model = RetinaFace(network_name='resnet50', device=device, half=half)
load_net = torch.load(model_path, map_location=lambda storage, loc: storage)
# remove unnecessary 'module.'
for k, v in deepcopy(load_net).items():
if k.startswith('module.'):
load_net[k[7:]] = v
load_net.pop(k)
model.load_state_dict(load_net, strict=True)
return model
def tensor_to_rgb(tensor_image):
"""
将ComfyUI的tensor图像转换为RGB图像
参数:
tensor_image: 形状为[B,H,W,C]的tensor,通常为float32类型,值范围0-1
返回:
numpy数组,RGB格式,uint8类型,值范围0-255
"""
# ComfyUI的tensor格式是[B,H,W,C],取第一张图片
if len(tensor_image.shape) == 4:
image = tensor_image[0].cpu().numpy()
else:
image = tensor_image.cpu().numpy()
# 转换为0-255范围的uint8
image = (image * 255.0).astype(np.uint8)
return image
def rgb_to_tensor(image):
"""
将RGB图像转换回ComfyUI的tensor格式
参数:
image: numpy数组,RGB格式,uint8类型
返回:
形状为[1,H,W,C]的tensor,float32类型,值范围0-1
"""
# 转换为float32并归一化到0-1
image = image.astype(np.float32) / 255.0
# 转换为tensor并添加批次维度
image = torch.from_numpy(image).unsqueeze(0)
return image
class AlignFace:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"half": ("BOOLEAN", {"default": False}),
# "unload_model": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "detect_and_align_whole_image"
CATEGORY = "🎤MW/MW-PortraitTools"
def detect_and_align_whole_image(self,
image,
half,
conf_threshold=0.8,
nms_threshold=0.4,
use_origin_size=True):
# 初始化模型
face_detector = init_model(half=half, device=device)
# 读取图像
img_rgb = tensor_to_rgb(image)
# 检测人脸
face_info = face_detector.detect_faces(img_rgb,
conf_threshold=conf_threshold,
nms_threshold=nms_threshold,
use_origin_size=use_origin_size
)
if len(face_info) == 0:
print("未检测到人脸")
return (image,)
# 获取最大的人脸(假设主要人脸是最大的)
areas = (face_info[:, 2] - face_info[:, 0]) * (face_info[:, 3] - face_info[:, 1])
max_face_idx = np.argmax(areas)
face_box = face_info[max_face_idx, 0:4]
landmarks = face_info[max_face_idx, 5:15].reshape(5, 2)
# 提取关键点
facial5points = [[landmarks[j][0], landmarks[j][1]] for j in range(5)]
# 使用简单的方法进行对齐
# 计算眼睛中心点(使用前两个关键点,它们通常是左右眼)
left_eye = facial5points[0]
right_eye = facial5points[1]
# 计算眼睛之间的角度
dy = right_eye[1] - left_eye[1]
dx = right_eye[0] - left_eye[0]
angle = np.degrees(np.arctan2(dy, dx))
# 计算眼睛中心
eye_center = ((left_eye[0] + right_eye[0]) // 2, (left_eye[1] + right_eye[1]) // 2)
# 获取旋转矩阵
M = cv2.getRotationMatrix2D(eye_center, angle, 1)
# 对整个图像进行旋转
h, w = img_rgb.shape[:2]
aligned_image = cv2.warpAffine(img_rgb, M, (w, h), flags=cv2.INTER_CUBIC, borderMode=cv2.BORDER_REPLICATE)
image_tensor = rgb_to_tensor(aligned_image)
# 重新检测旋转后的人脸,以获取新的边界框
rotated_face_info = face_detector.detect_faces(aligned_image,
conf_threshold=conf_threshold,
nms_threshold=nms_threshold,
use_origin_size=use_origin_size
)
if len(rotated_face_info) == 0:
print("对齐后未检测到人脸,返回原始图像")
return (image,)
else:
return (image_tensor,)
class DetectCropFaces:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"half": ("BOOLEAN", {"default": False}),
"horizontal_padding": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.1}),
"vertical_padding": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.1}),
"do_align": ("BOOLEAN", {"default": True}),
"angle_offset": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.1}),
# "unload_model": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "detect_and_align_faces"
CATEGORY = "🎤MW/MW-PortraitTools"
def detect_and_align_faces(self,
image,
half,
horizontal_padding,
vertical_padding,
do_align=True,
angle_offset=0.1,
conf_threshold=0.8,
nms_threshold=0.4,
use_origin_size=True):
# 初始化模型
face_detector = init_model(half=half, device=device)
# 读取图像
img_rgb = tensor_to_rgb(image)
_, aligned_faces = face_detector.align_multi(
img_rgb,
padding=(horizontal_padding, vertical_padding),
do_align=do_align,
angle_offset=angle_offset,
conf_threshold=conf_threshold,
nms_threshold=nms_threshold,
use_origin_size=use_origin_size,
limit=None)
if len(aligned_faces) == 0:
print("没有检测到人脸,返回原始图像")
return (image,)
# 如果只检测到一个人脸,直接返回
if len(aligned_faces) == 1:
face_tensor = rgb_to_tensor(aligned_faces[0])
return (face_tensor,)
face_tensors = []
# 找出所有人脸中最大的尺寸
max_h = max([face.shape[0] for face in aligned_faces])
max_w = max([face.shape[1] for face in aligned_faces])
for i, face in enumerate(aligned_faces):
# 调整所有人脸到相同大小
resized_face = cv2.resize(face, (max_w, max_h), interpolation=cv2.INTER_CUBIC)
face_tensor = rgb_to_tensor(resized_face)
face_tensors.append(face_tensor)
# 将所有人脸tensor拼接成一个批次
batch_tensor = torch.cat(face_tensors, dim=0)
# 返回批次tensor
return (batch_tensor,)
size_list = [
"一寸,413,295",
"二寸,626,413",
"小一寸,378,260",
"小二寸,531,413",
"大一寸,567,390",
"大二寸,626,413",
"五寸,1499,1050",
"教师资格证,413,295",
"国家公务员考试,413,295",
"初级会计考试,413,295",
"英语四六级考试,192,144",
"计算机等级考试,567,390",
"研究生考试,709,531",
"社保卡,441,358",
"电子驾驶证,378,260",
"美国签证,600,600",
"日本签证,413,295",
"韩国签证,531,413"
]
bg_colors = {
"Alpha": None,
"black": (0, 0, 0),
"white": (255, 255, 255),
"gray": (128, 128, 128),
"green": (0, 255, 0),
"pure_blue": (0, 0, 255),
"pure_red": (255, 0, 0),
"cornflower_blue": (98, 139, 206),
"crimson_red": (215, 69, 50),
"dark_slate_blue": (75, 97, 144),
"snow_white": (242, 240, 240)
}
class IDPhotos:
def __init__(self):
self.models = {
"RMBG-2.0": RMBGModel(),
"INSPYRENET": InspyrenetModel(),
"BEN": BENModel(),
"BEN2": BEN2Model()
}
@classmethod
def INPUT_TYPES(s):
return {
"required":{
"image":("IMAGE",),
"rmbg_model":(list(AVAILABLE_MODELS),{"default":"RMBG-2.0"}),
"bg_color":(list(bg_colors.keys()),{"default":"Alpha"}),
"size":(size_list,{"default":"一寸,413,295"}),
"kb":("INT",{"default":500,"min":5,"max":2000,"step":1}),
"dpi":("INT",{"default":300,"min":50,"max":1000,"step":10}),
"face_reduction":("FLOAT",{
"default": 1.0,
"min":0.0,
"max":5.0,
"step":0.1,
}),
"face_up_down":("FLOAT",{
"default": 0.0,
"min":-0.5,
"max":0.5,
"step":0.01,
}),
"angle_offset":("FLOAT",{
"default": 1.0,
"min":-10.0,
"max":10.0,
"step":0.1,
}),
}
}
RETURN_TYPES = ("IMAGE", "IMAGE", "IMAGE")
RETURN_NAMES = ("standard_photo", "hd_photo", "print_photos")
FUNCTION = "gen_img"
CATEGORY = "🎤MW/MW-PortraitTools"
def gen_img(self, image, rmbg_model, bg_color, size, face_reduction, face_up_down, angle_offset, kb, dpi=300):
# 解析尺寸参数
size_parts = size.split(',')
size = (int(size_parts[1]), int(size_parts[2]))
hd_photo = self.photo_gen(image, size, face_reduction, face_up_down, angle_offset, by_size=False)
rmbg_hd_photo = self.image_rmbg(hd_photo, rmbg_model, bg_color)
standard_photo = self.photo_gen(image, size, face_reduction, face_up_down, angle_offset, by_size=True)
rmbg_standard_photo = self.image_rmbg(standard_photo, rmbg_model, bg_color)
print_photos = self.print_photos_gen(rmbg_standard_photo, size, kb, dpi)
return (rmbg_standard_photo, rmbg_hd_photo, print_photos)
def photo_gen(self, image, size, face_reduction, face_up_down, angle_offset, by_size=True):
# 初始化模型
face_detector = init_model(half=True, device=device)
# 读取图像
img_rgb = tensor_to_rgb(image)
_, aligned_faces = face_detector.align_multi(
img_rgb,
angle_offset=angle_offset,
padding=(face_reduction, face_reduction),
do_align=True,
conf_threshold=0.8,
nms_threshold=0.4,
use_origin_size=True,
limit=None)
if len(aligned_faces) > 1:
raise ValueError("Multiple faces detected, please upload an image of a single face.")
if len(aligned_faces) == 0:
raise ValueError("No face detected, please upload an image containing the face.")
face = aligned_faces[0]
# 获取目标尺寸
target_h = size[0]
target_w = size[1]
# 获取人脸图像的尺寸
face_h, face_w = face.shape[:2]
if by_size:
# 模式1: 按指定尺寸调整
# 计算缩放比例,使照片高度比目标高度多 target_h * face_up_down
scale_h = (target_h + target_h * abs(face_up_down)) / face_h
scaled_w = int(face_w * scale_h)
scaled_h = int(face_h * scale_h)
# 如果宽度小于目标宽度,继续放大
if scaled_w < target_w:
scale_w = target_w / scaled_w
scaled_w = target_w
scaled_h = int(scaled_h * scale_w)
# 缩放图像
scaled_face = cv2.resize(face, (scaled_w, scaled_h), interpolation=cv2.INTER_LANCZOS4)
# 计算裁剪区域
center_x = scaled_w // 2
# 移动裁剪区域
center_y = int(scaled_h // 2 + target_h * face_up_down / 2)
# 计算裁剪区域的左上角和右下角坐标
left = center_x - target_w // 2
right = left + target_w
top = center_y - target_h // 2
bottom = top + target_h
# 确保裁剪区域在图像内
if left < 0:
left = 0
right = target_w
if right > scaled_w:
right = scaled_w
left = scaled_w - target_w
if top < 0:
top = 0
bottom = target_h
if bottom > scaled_h:
bottom = scaled_h
top = scaled_h - target_h
# 裁剪图像
final_image = scaled_face[top:bottom, left:right]
# 确保最终图像尺寸正确
if final_image.shape[:2] != (target_h, target_w):
final_image = cv2.resize(final_image, (target_w, target_h), interpolation=cv2.INTER_LANCZOS4)
else:
# 模式2: 严格按比例调整尺寸,不缩放照片
# 计算缩放比例,使照片高度比目标高度多 target_h * face_up_down
adjusted_target_h = int(face_h/(1 + abs(face_up_down)))
adjusted_target_w = int(target_w * (adjusted_target_h / target_h))
# 如果调整后的宽度超过了照片宽度,需要重新计算
if adjusted_target_w > face_w:
# 按照宽度计算比例
scale_w = face_w / target_w
adjusted_target_w = face_w
adjusted_target_h = int(target_h * scale_w)
# 计算裁剪区域
center_x = face_w // 2
# 移动裁剪区域
center_y = int(face_h // 2 + adjusted_target_h * face_up_down / 2)
# 计算裁剪区域的左上角和右下角坐标
left = center_x - adjusted_target_w // 2
right = left + adjusted_target_w
top = center_y - adjusted_target_h // 2
bottom = top + adjusted_target_h
# 确保裁剪区域在图像内
if left < 0:
left = 0
right = adjusted_target_w
if right > face_w:
right = face_w
left = face_w - adjusted_target_w
if top < 0:
top = 0
bottom = adjusted_target_h
if bottom > face_h:
bottom = face_h
top = face_h - adjusted_target_h
# 裁剪图像
final_image = face[top:bottom, left:right]
# 转换回tensor格式
result_tensor = rgb_to_tensor(final_image)
return result_tensor
def print_photos_gen(self, input_image, size, kb, dpi=300):
# 将tensor转换为RGB图像
img_rgb = tensor_to_rgb(input_image)
from io import BytesIO
pil_img = Image.fromarray(img_rgb)
# 创建字节流对象
img_byte_arr = BytesIO()
# 保存到字节流
pil_img.save(img_byte_arr, format="PNG", dpi=(dpi, dpi))
img_byte_arr.seek(0)
# 调整图像大小到指定KB
quality = 95
while True:
# 创建字节流对象
img_byte_arr = BytesIO()
# 保存图像到字节流
pil_img.save(img_byte_arr, format="PNG", quality=quality, dpi=(dpi, dpi))
# 获取图像大小(KB)
img_size_kb = len(img_byte_arr.getvalue()) / 1024
# 检查图像大小是否在目标范围内
if img_size_kb <= kb or quality == 1:
# 如果图像小于目标大小,添加填充
if img_size_kb < kb:
padding_size = int(
(kb * 1024) - len(img_byte_arr.getvalue())
)
padding = b"\x00" * padding_size
img_byte_arr.write(padding)
break
# 如果图像仍然太大,降低质量
quality -= 5
# 确保质量不低于1
if quality < 1:
quality = 1
# 将字节流转换回PIL图像
img_byte_arr.seek(0)
pil_img = Image.open(img_byte_arr)
result_layout_photo = cv2.cvtColor(np.array(pil_img), cv2.COLOR_BGR2RGB)
# 生成布局
typography_arr, typography_rotate = generate_layout_photo(
input_height=size[0], input_width=size[1]
)
# 生成最终布局图像
result_layout_image = generate_layout_image(
result_layout_photo,
typography_arr,
typography_rotate,
height=size[0],
width=size[1],
)
# 转换为RGB并转换为tensor
print_cv2 = cv2.cvtColor(result_layout_image, cv2.COLOR_BGR2RGB)
print_photos = rgb_to_tensor(print_cv2)
return print_photos
def image_rmbg(self, image, model, bg_color):
model_instance = self.models[model]
params = {
"sensitivity": 1.0,
"process_res": 1024,
"mask_blur": 0,
"mask_offset": 0,
"background": bg_colors[bg_color],
"invert_output": False,
"optimize": "default",
"refine_foreground": False
}
# Check and download model if needed
cache_status, message = model_instance.check_model_cache(model)
if not cache_status:
print(f"Cache check: {message}")
print("Downloading required model files...")
download_status, download_message = model_instance.download_model(model)
if not download_status:
handle_model_error(download_message)
print("Model files downloaded successfully")
# Get mask from specific model
mask = model_instance.process_image(image, model, params)
# Ensure mask is in the correct format
if isinstance(mask, list):
masks = [m.convert("L") for m in mask if isinstance(m, Image.Image)]
mask = masks[0] if masks else None
elif isinstance(mask, Image.Image):
mask = mask.convert("L")
# Post-process mask
mask_tensor = pil2tensor(mask)
mask_tensor = mask_tensor * (1 + (1 - params["sensitivity"]))
mask_tensor = torch.clamp(mask_tensor, 0, 1)
mask = tensor2pil(mask_tensor)
# Create final image
orig_image = tensor2pil(image)
orig_rgba = orig_image.convert("RGBA")
r, g, b, _ = orig_rgba.split()
foreground = Image.merge('RGBA', (r, g, b, mask))
if bg_color != "Alpha":
bg_color = bg_colors[bg_color]
bg_image = Image.new('RGBA', orig_image.size, (*bg_color, 255))
composite_image = Image.alpha_composite(bg_image, foreground)
processed_image = pil2tensor(composite_image.convert("RGB"))
else:
processed_image = pil2tensor(foreground)
return processed_image
class BeautifyPhoto:
@classmethod
def INPUT_TYPES(s):
return {
"required":{
"image":("IMAGE",),
"whitening_strength":("INT",{
"default": 0,
"min":0,
"max":100,
"step":1,
}),
"brightness_strength":("INT",{
"default": 0,
"min":-100,
"max":100,
"step":1,
}),
"contrast_strength":("INT",{
"default": 0,
"min":-100,
"max":100,
"step":1,
}),
"saturation_strength":("INT",{
"default": 0,
"min":-100,
"max":100,
"step":1,
}),
"sharpen_strength":("FLOAT",{
"default": 0.1,
"min":0.0,
"max":10.0,
"step":0.1,
}),
"grind_skin":("INT",{
"default": 0,
"min":0,
"max":10,
"step":1,
}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "beautify"
CATEGORY = "🎤MW/MW-PortraitTools"
def beautify(self,
image,
whitening_strength,
brightness_strength,
contrast_strength,
saturation_strength,
sharpen_strength,
grind_skin,
):
img_np = tensor_to_rgb(image)
input_image = cv2.cvtColor(img_np,cv2.COLOR_BGR2RGB)
adjusted_image = grindSkin(input_image, strength=grind_skin)
adjusted_image = make_whitening(adjusted_image, strength=whitening_strength)
adjusted_image = adjust_brightness_contrast_sharpen_saturation(
adjusted_image,
brightness_factor=brightness_strength,
contrast_factor=contrast_strength,
sharpen_strength=sharpen_strength,
saturation_factor=saturation_strength,
)
result_image = cv2.cvtColor(adjusted_image,cv2.COLOR_BGR2RGB)
result_tensor = rgb_to_tensor(result_image)
return (result_tensor,)
NODE_CLASS_MAPPINGS = {
"DetectCropFace": DetectCropFaces,
"AlignFace": AlignFace,
"IDPhotos": IDPhotos,
"BeautifyPhoto": BeautifyPhoto,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DetectCropFaces": "Detect and Crop Faces",
"AlignFace": "Align Face",
"IDPhotos": "ID Photos",
"BeautifyPhoto": "Beautify Photo",
}
+15
View File
@@ -0,0 +1,15 @@
[project]
name = "audiotools-mw"
description = "Portrait Tools: Facial detection cropping, alignment, ID photo, etc."
version = "1.0.0"
license = {file = "LICENSE"}
dependencies = []
[project.urls]
Repository = "https://github.com/billwuhao/ComfyUI_PortraitTools"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "mw"
DisplayName = "MW-ComfyUI_PortraitTools"
Icon = ""
+2
View File
@@ -0,0 +1,2 @@
huggingface_hub
opencv-python
+1
View File
@@ -0,0 +1 @@
from retinaface.retinaface import RetinaFace
+479
View File
@@ -0,0 +1,479 @@
import cv2
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from PIL import Image
from torchvision.models._utils import IntermediateLayerGetter as IntermediateLayerGetter
from retinaface.retinaface_net import FPN, SSH, MobileNetV1, make_bbox_head, make_class_head, make_landmark_head
from retinaface.retinaface_utils import (PriorBox, batched_decode, batched_decode_landm, decode, decode_landm,
py_cpu_nms)
def generate_config(network_name):
cfg_re50 = {
'name': 'Resnet50',
'min_sizes': [[16, 32], [64, 128], [256, 512]],
'steps': [8, 16, 32],
'variance': [0.1, 0.2],
'clip': False,
'loc_weight': 2.0,
'gpu_train': True,
'batch_size': 24,
'ngpu': 4,
'epoch': 100,
'decay1': 70,
'decay2': 90,
'image_size': 840,
'return_layers': {
'layer2': 1,
'layer3': 2,
'layer4': 3
},
'in_channel': 256,
'out_channel': 256
}
if network_name == 'resnet50':
return cfg_re50
else:
raise NotImplementedError(f'network_name={network_name}')
class RetinaFace(nn.Module):
def __init__(self, network_name='resnet50', device=torch.device('cuda'), half=False, phase='test'):
super(RetinaFace, self).__init__()
self.half_inference = half
cfg = generate_config(network_name)
self.backbone = cfg['name']
self.device = device
self.cfg = cfg
self.phase = phase
self.target_size, self.max_size = 1600, 2150
self.resize, self.scale, self.scale1 = 1., None, None
self.mean_tensor = torch.tensor([[[[104.]], [[117.]], [[123.]]]]).to(self.device)
# Build network.
backbone = None
if cfg['name'] == 'mobilenet0.25':
backbone = MobileNetV1()
self.body = IntermediateLayerGetter(backbone, cfg['return_layers'])
elif cfg['name'] == 'Resnet50':
import torchvision.models as models
backbone = models.resnet50(pretrained=False)
self.body = IntermediateLayerGetter(backbone, cfg['return_layers'])
in_channels_stage2 = cfg['in_channel']
in_channels_list = [
in_channels_stage2 * 2,
in_channels_stage2 * 4,
in_channels_stage2 * 8,
]
out_channels = cfg['out_channel']
self.fpn = FPN(in_channels_list, out_channels)
self.ssh1 = SSH(out_channels, out_channels)
self.ssh2 = SSH(out_channels, out_channels)
self.ssh3 = SSH(out_channels, out_channels)
self.ClassHead = make_class_head(fpn_num=3, inchannels=cfg['out_channel'])
self.BboxHead = make_bbox_head(fpn_num=3, inchannels=cfg['out_channel'])
self.LandmarkHead = make_landmark_head(fpn_num=3, inchannels=cfg['out_channel'])
self.to(self.device)
self.eval()
if self.half_inference:
self.half()
def forward(self, inputs):
out = self.body(inputs)
if self.backbone == 'mobilenet0.25' or self.backbone == 'Resnet50':
out = list(out.values())
# FPN
fpn = self.fpn(out)
# SSH
feature1 = self.ssh1(fpn[0])
feature2 = self.ssh2(fpn[1])
feature3 = self.ssh3(fpn[2])
features = [feature1, feature2, feature3]
bbox_regressions = torch.cat([self.BboxHead[i](feature) for i, feature in enumerate(features)], dim=1)
classifications = torch.cat([self.ClassHead[i](feature) for i, feature in enumerate(features)], dim=1)
tmp = [self.LandmarkHead[i](feature) for i, feature in enumerate(features)]
ldm_regressions = (torch.cat(tmp, dim=1))
if self.phase == 'train':
output = (bbox_regressions, classifications, ldm_regressions)
else:
output = (bbox_regressions, F.softmax(classifications, dim=-1), ldm_regressions)
return output
def __detect_faces(self, inputs):
# get scale
height, width = inputs.shape[2:]
self.scale = torch.tensor([width, height, width, height], dtype=torch.float32).to(self.device)
tmp = [width, height, width, height, width, height, width, height, width, height]
self.scale1 = torch.tensor(tmp, dtype=torch.float32).to(self.device)
# forawrd
inputs = inputs.to(self.device)
if self.half_inference:
inputs = inputs.half()
loc, conf, landmarks = self(inputs)
# get priorbox
priorbox = PriorBox(self.cfg, image_size=inputs.shape[2:])
priors = priorbox.forward().to(self.device)
return loc, conf, landmarks, priors
# single image detection
def transform(self, image, use_origin_size):
# convert to opencv format
if isinstance(image, Image.Image):
image = cv2.cvtColor(np.asarray(image), cv2.COLOR_RGB2BGR)
image = image.astype(np.float32)
# testing scale
im_size_min = np.min(image.shape[0:2])
im_size_max = np.max(image.shape[0:2])
resize = float(self.target_size) / float(im_size_min)
# prevent bigger axis from being more than max_size
if np.round(resize * im_size_max) > self.max_size:
resize = float(self.max_size) / float(im_size_max)
resize = 1 if use_origin_size else resize
# resize
if resize != 1:
image = cv2.resize(image, None, None, fx=resize, fy=resize, interpolation=cv2.INTER_LINEAR)
# convert to torch.tensor format
# image -= (104, 117, 123)
image = image.transpose(2, 0, 1)
image = torch.from_numpy(image).unsqueeze(0)
return image, resize
def detect_faces(
self,
image,
conf_threshold=0.8,
nms_threshold=0.4,
use_origin_size=True,
):
"""
Params:
imgs: BGR image
"""
image, self.resize = self.transform(image, use_origin_size)
image = image.to(self.device)
if self.half_inference:
image = image.half()
image = image - self.mean_tensor
loc, conf, landmarks, priors = self.__detect_faces(image)
boxes = decode(loc.data.squeeze(0), priors.data, self.cfg['variance'])
boxes = boxes * self.scale / self.resize
boxes = boxes.cpu().numpy()
scores = conf.squeeze(0).data.cpu().numpy()[:, 1]
landmarks = decode_landm(landmarks.squeeze(0), priors, self.cfg['variance'])
landmarks = landmarks * self.scale1 / self.resize
landmarks = landmarks.detach().cpu().numpy()
# ignore low scores
inds = np.where(scores > conf_threshold)[0]
boxes, landmarks, scores = boxes[inds], landmarks[inds], scores[inds]
# sort
order = scores.argsort()[::-1]
boxes, landmarks, scores = boxes[order], landmarks[order], scores[order]
# do NMS
bounding_boxes = np.hstack((boxes, scores[:, np.newaxis])).astype(np.float32, copy=False)
keep = py_cpu_nms(bounding_boxes, nms_threshold)
bounding_boxes, landmarks = bounding_boxes[keep, :], landmarks[keep]
# self.t['forward_pass'].toc()
# print(self.t['forward_pass'].average_time)
# import sys
# sys.stdout.flush()
return np.concatenate((bounding_boxes, landmarks), axis=1)
def __align_multi(self, image, boxes, landmarks, angle_offset=0.1, padding=(0.5, 0.4), do_align=True, limit=None):
"""
对检测到的多个人脸进行对齐和裁剪
参数:
image: 输入图像
boxes: 人脸边界框
landmarks: 人脸关键点
limit: 处理的人脸数量限制
返回:
人脸信息和对齐后的人脸图像
"""
if len(boxes) < 1:
return [], []
if limit:
boxes = boxes[:limit]
landmarks = landmarks[:limit]
faces = []
for i, landmark in enumerate(landmarks):
facial5points = [[landmark[2 * j], landmark[2 * j + 1]] for j in range(5)]
# 计算原始人脸区域的大小
x1, y1, x2, y2, _ = boxes[i]
face_width = int(x2 - x1)
face_height = int(y2 - y1)
# 增加裁剪区域,确保包含下巴和耳朵
horizontal_padding = int(face_width * padding[0])
vertical_padding = int(face_height * padding[1])
# 计算扩展后的裁剪尺寸
crop_width = face_width + 2 * horizontal_padding
crop_height = face_height + 2 * vertical_padding
# 确保尺寸为偶数,便于处理
crop_width = crop_width if crop_width % 2 == 0 else crop_width + 1
crop_height = crop_height if crop_height % 2 == 0 else crop_height + 1
if do_align:
try:
# 计算眼睛位置(通常前两个关键点是左右眼)
left_eye = facial5points[0]
right_eye = facial5points[1]
# 计算眼睛之间的角度
dy = right_eye[1] - left_eye[1]
dx = right_eye[0] - left_eye[0]
angle = np.degrees(np.arctan2(dy, dx))
# 应用角度微调
angle += angle_offset
# 计算眼睛中心
eye_center = ((left_eye[0] + right_eye[0]) / 2, (left_eye[1] + right_eye[1]) / 2)
# 创建一个足够大的画布,确保旋转后不会裁剪到人脸
h, w = image.shape[:2]
diagonal = int(np.sqrt(crop_width**2 + crop_height**2)) + 20 # 额外添加一些边距
canvas_size = (diagonal, diagonal)
# 计算人脸中心
face_center_x = (x1 + x2) / 2
face_center_y = (y1 + y2) / 2
# 计算偏移,使人脸在画布中居中
offset_x = int(diagonal / 2 - face_center_x)
offset_y = int(diagonal / 2 - face_center_y)
# 创建画布并将原图放在适当位置
canvas = np.zeros((diagonal, diagonal, 3), dtype=np.uint8)
# 计算原图在画布上的位置
src_x1 = max(0, -offset_x)
src_y1 = max(0, -offset_y)
src_x2 = min(w, diagonal - offset_x)
src_y2 = min(h, diagonal - offset_y)
dst_x1 = max(0, offset_x)
dst_y1 = max(0, offset_y)
dst_x2 = min(diagonal, w + offset_x)
dst_y2 = min(diagonal, h + offset_y)
canvas[dst_y1:dst_y2, dst_x1:dst_x2] = image[src_y1:src_y2, src_x1:src_x2]
# 调整眼睛中心坐标
adjusted_eye_center = (eye_center[0] + offset_x, eye_center[1] + offset_y)
# 对画布进行旋转
M = cv2.getRotationMatrix2D(adjusted_eye_center, angle, 1)
rotated_canvas = cv2.warpAffine(canvas, M, canvas_size, flags=cv2.INTER_CUBIC, borderMode=cv2.BORDER_REPLICATE)
# 计算旋转后的人脸中心
rotated_center = np.dot(M, [adjusted_eye_center[0], adjusted_eye_center[1], 1])
# 计算裁剪区域,确保人脸居中
crop_x1 = int(rotated_center[0] - crop_width / 2)
crop_y1 = int(rotated_center[1] - crop_height / 2)
# 确保裁剪区域在画布内
crop_x1 = max(0, min(crop_x1, diagonal - crop_width))
crop_y1 = max(0, min(crop_y1, diagonal - crop_height))
crop_x2 = crop_x1 + crop_width
crop_y2 = crop_y1 + crop_height
# 裁剪
face_img = rotated_canvas[crop_y1:crop_y2, crop_x1:crop_x2]
# 如果裁剪后的尺寸与预期不符,调整大小
if face_img.shape[0] != crop_height or face_img.shape[1] != crop_width:
face_img = cv2.resize(face_img, (crop_width, crop_height))
faces.append(face_img)
except Exception as e:
print(f"对齐人脸时出错: {e}")
# 出错时,直接裁剪原始人脸区域作为备选,并增加边距
try:
# 计算扩展后的边界框
ext_x1 = max(0, int(x1 - horizontal_padding))
ext_y1 = max(0, int(y1 - vertical_padding))
ext_x2 = min(image.shape[1], int(x2 + horizontal_padding))
ext_y2 = min(image.shape[0], int(y2 + vertical_padding))
# 直接裁剪
face_img = image[ext_y1:ext_y2, ext_x1:ext_x2]
# 调整大小
face_img = cv2.resize(face_img, (crop_width, crop_height))
faces.append(face_img)
except Exception as e2:
print(f"裁剪人脸备选方案也失败: {e2}")
# 如果都失败了,添加一个空白图像
faces.append(np.zeros((crop_height, crop_width, 3), dtype=np.uint8))
else:
# 直接裁剪原始人脸区域
ext_x1 = max(0, int(x1 - horizontal_padding))
ext_y1 = max(0, int(y1 - vertical_padding))
ext_x2 = min(image.shape[1], int(x2 + horizontal_padding))
ext_y2 = min(image.shape[0], int(y2 + vertical_padding))
face_img = image[ext_y1:ext_y2, ext_x1:ext_x2]
face_img = cv2.resize(face_img, (crop_width, crop_height))
faces.append(face_img)
return np.concatenate((boxes, landmarks), axis=1), faces
def align_multi(self,
img,
angle_offset=0.1,
padding=(0.5, 0.4),
do_align=True,
conf_threshold=0.8,
nms_threshold=0.4,
use_origin_size=True,
limit=None):
rlt = self.detect_faces(img, conf_threshold=conf_threshold,
nms_threshold=nms_threshold,
use_origin_size=use_origin_size,)
boxes, landmarks = rlt[:, 0:5], rlt[:, 5:]
return self.__align_multi(img, boxes, landmarks, angle_offset=angle_offset, padding=padding, do_align=do_align, limit=limit)
# batched detection
def batched_transform(self, frames, use_origin_size):
"""
Arguments:
frames: a list of PIL.Image, or torch.Tensor(shape=[n, h, w, c],
type=np.float32, BGR format).
use_origin_size: whether to use origin size.
"""
from_PIL = True if isinstance(frames[0], Image.Image) else False
# convert to opencv format
if from_PIL:
frames = [cv2.cvtColor(np.asarray(frame), cv2.COLOR_RGB2BGR) for frame in frames]
frames = np.asarray(frames, dtype=np.float32)
# testing scale
im_size_min = np.min(frames[0].shape[0:2])
im_size_max = np.max(frames[0].shape[0:2])
resize = float(self.target_size) / float(im_size_min)
# prevent bigger axis from being more than max_size
if np.round(resize * im_size_max) > self.max_size:
resize = float(self.max_size) / float(im_size_max)
resize = 1 if use_origin_size else resize
# resize
if resize != 1:
if not from_PIL:
frames = F.interpolate(frames, scale_factor=resize)
else:
frames = [
cv2.resize(frame, None, None, fx=resize, fy=resize, interpolation=cv2.INTER_LINEAR)
for frame in frames
]
# convert to torch.tensor format
if not from_PIL:
frames = frames.transpose(1, 2).transpose(1, 3).contiguous()
else:
frames = frames.transpose((0, 3, 1, 2))
frames = torch.from_numpy(frames)
return frames, resize
def batched_detect_faces(self, frames, conf_threshold=0.8, nms_threshold=0.4, use_origin_size=True):
"""
Arguments:
frames: a list of PIL.Image, or np.array(shape=[n, h, w, c],
type=np.uint8, BGR format).
conf_threshold: confidence threshold.
nms_threshold: nms threshold.
use_origin_size: whether to use origin size.
Returns:
final_bounding_boxes: list of np.array ([n_boxes, 5],
type=np.float32).
final_landmarks: list of np.array ([n_boxes, 10], type=np.float32).
"""
# self.t['forward_pass'].tic()
frames, self.resize = self.batched_transform(frames, use_origin_size)
frames = frames.to(self.device)
frames = frames - self.mean_tensor
b_loc, b_conf, b_landmarks, priors = self.__detect_faces(frames)
final_bounding_boxes, final_landmarks = [], []
# decode
priors = priors.unsqueeze(0)
b_loc = batched_decode(b_loc, priors, self.cfg['variance']) * self.scale / self.resize
b_landmarks = batched_decode_landm(b_landmarks, priors, self.cfg['variance']) * self.scale1 / self.resize
b_conf = b_conf[:, :, 1]
# index for selection
b_indice = b_conf > conf_threshold
# concat
b_loc_and_conf = torch.cat((b_loc, b_conf.unsqueeze(-1)), dim=2).float()
for pred, landm, inds in zip(b_loc_and_conf, b_landmarks, b_indice):
# ignore low scores
pred, landm = pred[inds, :], landm[inds, :]
if pred.shape[0] == 0:
final_bounding_boxes.append(np.array([], dtype=np.float32))
final_landmarks.append(np.array([], dtype=np.float32))
continue
# sort
# order = score.argsort(descending=True)
# box, landm, score = box[order], landm[order], score[order]
# to CPU
bounding_boxes, landm = pred.detach().cpu().numpy(), landm.detach().cpu().numpy()
# NMS
keep = py_cpu_nms(bounding_boxes, nms_threshold)
bounding_boxes, landmarks = bounding_boxes[keep, :], landm[keep]
# append
final_bounding_boxes.append(bounding_boxes)
final_landmarks.append(landmarks)
# self.t['forward_pass'].toc(average=True)
# self.batch_time += self.t['forward_pass'].diff
# self.total_frame += len(frames)
# print(self.batch_time / self.total_frame)
return final_bounding_boxes, final_landmarks
+196
View File
@@ -0,0 +1,196 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
def conv_bn(inp, oup, stride=1, leaky=0):
return nn.Sequential(
nn.Conv2d(inp, oup, 3, stride, 1, bias=False), nn.BatchNorm2d(oup),
nn.LeakyReLU(negative_slope=leaky, inplace=True))
def conv_bn_no_relu(inp, oup, stride):
return nn.Sequential(
nn.Conv2d(inp, oup, 3, stride, 1, bias=False),
nn.BatchNorm2d(oup),
)
def conv_bn1X1(inp, oup, stride, leaky=0):
return nn.Sequential(
nn.Conv2d(inp, oup, 1, stride, padding=0, bias=False), nn.BatchNorm2d(oup),
nn.LeakyReLU(negative_slope=leaky, inplace=True))
def conv_dw(inp, oup, stride, leaky=0.1):
return nn.Sequential(
nn.Conv2d(inp, inp, 3, stride, 1, groups=inp, bias=False),
nn.BatchNorm2d(inp),
nn.LeakyReLU(negative_slope=leaky, inplace=True),
nn.Conv2d(inp, oup, 1, 1, 0, bias=False),
nn.BatchNorm2d(oup),
nn.LeakyReLU(negative_slope=leaky, inplace=True),
)
class SSH(nn.Module):
def __init__(self, in_channel, out_channel):
super(SSH, self).__init__()
assert out_channel % 4 == 0
leaky = 0
if (out_channel <= 64):
leaky = 0.1
self.conv3X3 = conv_bn_no_relu(in_channel, out_channel // 2, stride=1)
self.conv5X5_1 = conv_bn(in_channel, out_channel // 4, stride=1, leaky=leaky)
self.conv5X5_2 = conv_bn_no_relu(out_channel // 4, out_channel // 4, stride=1)
self.conv7X7_2 = conv_bn(out_channel // 4, out_channel // 4, stride=1, leaky=leaky)
self.conv7x7_3 = conv_bn_no_relu(out_channel // 4, out_channel // 4, stride=1)
def forward(self, input):
conv3X3 = self.conv3X3(input)
conv5X5_1 = self.conv5X5_1(input)
conv5X5 = self.conv5X5_2(conv5X5_1)
conv7X7_2 = self.conv7X7_2(conv5X5_1)
conv7X7 = self.conv7x7_3(conv7X7_2)
out = torch.cat([conv3X3, conv5X5, conv7X7], dim=1)
out = F.relu(out)
return out
class FPN(nn.Module):
def __init__(self, in_channels_list, out_channels):
super(FPN, self).__init__()
leaky = 0
if (out_channels <= 64):
leaky = 0.1
self.output1 = conv_bn1X1(in_channels_list[0], out_channels, stride=1, leaky=leaky)
self.output2 = conv_bn1X1(in_channels_list[1], out_channels, stride=1, leaky=leaky)
self.output3 = conv_bn1X1(in_channels_list[2], out_channels, stride=1, leaky=leaky)
self.merge1 = conv_bn(out_channels, out_channels, leaky=leaky)
self.merge2 = conv_bn(out_channels, out_channels, leaky=leaky)
def forward(self, input):
# names = list(input.keys())
# input = list(input.values())
output1 = self.output1(input[0])
output2 = self.output2(input[1])
output3 = self.output3(input[2])
up3 = F.interpolate(output3, size=[output2.size(2), output2.size(3)], mode='nearest')
output2 = output2 + up3
output2 = self.merge2(output2)
up2 = F.interpolate(output2, size=[output1.size(2), output1.size(3)], mode='nearest')
output1 = output1 + up2
output1 = self.merge1(output1)
out = [output1, output2, output3]
return out
class MobileNetV1(nn.Module):
def __init__(self):
super(MobileNetV1, self).__init__()
self.stage1 = nn.Sequential(
conv_bn(3, 8, 2, leaky=0.1), # 3
conv_dw(8, 16, 1), # 7
conv_dw(16, 32, 2), # 11
conv_dw(32, 32, 1), # 19
conv_dw(32, 64, 2), # 27
conv_dw(64, 64, 1), # 43
)
self.stage2 = nn.Sequential(
conv_dw(64, 128, 2), # 43 + 16 = 59
conv_dw(128, 128, 1), # 59 + 32 = 91
conv_dw(128, 128, 1), # 91 + 32 = 123
conv_dw(128, 128, 1), # 123 + 32 = 155
conv_dw(128, 128, 1), # 155 + 32 = 187
conv_dw(128, 128, 1), # 187 + 32 = 219
)
self.stage3 = nn.Sequential(
conv_dw(128, 256, 2), # 219 +3 2 = 241
conv_dw(256, 256, 1), # 241 + 64 = 301
)
self.avg = nn.AdaptiveAvgPool2d((1, 1))
self.fc = nn.Linear(256, 1000)
def forward(self, x):
x = self.stage1(x)
x = self.stage2(x)
x = self.stage3(x)
x = self.avg(x)
# x = self.model(x)
x = x.view(-1, 256)
x = self.fc(x)
return x
class ClassHead(nn.Module):
def __init__(self, inchannels=512, num_anchors=3):
super(ClassHead, self).__init__()
self.num_anchors = num_anchors
self.conv1x1 = nn.Conv2d(inchannels, self.num_anchors * 2, kernel_size=(1, 1), stride=1, padding=0)
def forward(self, x):
out = self.conv1x1(x)
out = out.permute(0, 2, 3, 1).contiguous()
return out.view(out.shape[0], -1, 2)
class BboxHead(nn.Module):
def __init__(self, inchannels=512, num_anchors=3):
super(BboxHead, self).__init__()
self.conv1x1 = nn.Conv2d(inchannels, num_anchors * 4, kernel_size=(1, 1), stride=1, padding=0)
def forward(self, x):
out = self.conv1x1(x)
out = out.permute(0, 2, 3, 1).contiguous()
return out.view(out.shape[0], -1, 4)
class LandmarkHead(nn.Module):
def __init__(self, inchannels=512, num_anchors=3):
super(LandmarkHead, self).__init__()
self.conv1x1 = nn.Conv2d(inchannels, num_anchors * 10, kernel_size=(1, 1), stride=1, padding=0)
def forward(self, x):
out = self.conv1x1(x)
out = out.permute(0, 2, 3, 1).contiguous()
return out.view(out.shape[0], -1, 10)
def make_class_head(fpn_num=3, inchannels=64, anchor_num=2):
classhead = nn.ModuleList()
for i in range(fpn_num):
classhead.append(ClassHead(inchannels, anchor_num))
return classhead
def make_bbox_head(fpn_num=3, inchannels=64, anchor_num=2):
bboxhead = nn.ModuleList()
for i in range(fpn_num):
bboxhead.append(BboxHead(inchannels, anchor_num))
return bboxhead
def make_landmark_head(fpn_num=3, inchannels=64, anchor_num=2):
landmarkhead = nn.ModuleList()
for i in range(fpn_num):
landmarkhead.append(LandmarkHead(inchannels, anchor_num))
return landmarkhead
+421
View File
@@ -0,0 +1,421 @@
import numpy as np
import torch
import torchvision
from itertools import product as product
from math import ceil
class PriorBox(object):
def __init__(self, cfg, image_size=None, phase='train'):
super(PriorBox, self).__init__()
self.min_sizes = cfg['min_sizes']
self.steps = cfg['steps']
self.clip = cfg['clip']
self.image_size = image_size
self.feature_maps = [[ceil(self.image_size[0] / step), ceil(self.image_size[1] / step)] for step in self.steps]
self.name = 's'
def forward(self):
anchors = []
for k, f in enumerate(self.feature_maps):
min_sizes = self.min_sizes[k]
for i, j in product(range(f[0]), range(f[1])):
for min_size in min_sizes:
s_kx = min_size / self.image_size[1]
s_ky = min_size / self.image_size[0]
dense_cx = [x * self.steps[k] / self.image_size[1] for x in [j + 0.5]]
dense_cy = [y * self.steps[k] / self.image_size[0] for y in [i + 0.5]]
for cy, cx in product(dense_cy, dense_cx):
anchors += [cx, cy, s_kx, s_ky]
# back to torch land
output = torch.Tensor(anchors).view(-1, 4)
if self.clip:
output.clamp_(max=1, min=0)
return output
def py_cpu_nms(dets, thresh):
"""Pure Python NMS baseline."""
keep = torchvision.ops.nms(
boxes=torch.Tensor(dets[:, :4]),
scores=torch.Tensor(dets[:, 4]),
iou_threshold=thresh,
)
return list(keep)
def point_form(boxes):
""" Convert prior_boxes to (xmin, ymin, xmax, ymax)
representation for comparison to point form ground truth data.
Args:
boxes: (tensor) center-size default boxes from priorbox layers.
Return:
boxes: (tensor) Converted xmin, ymin, xmax, ymax form of boxes.
"""
return torch.cat(
(
boxes[:, :2] - boxes[:, 2:] / 2, # xmin, ymin
boxes[:, :2] + boxes[:, 2:] / 2),
1) # xmax, ymax
def center_size(boxes):
""" Convert prior_boxes to (cx, cy, w, h)
representation for comparison to center-size form ground truth data.
Args:
boxes: (tensor) point_form boxes
Return:
boxes: (tensor) Converted xmin, ymin, xmax, ymax form of boxes.
"""
return torch.cat(
(boxes[:, 2:] + boxes[:, :2]) / 2, # cx, cy
boxes[:, 2:] - boxes[:, :2],
1) # w, h
def intersect(box_a, box_b):
""" We resize both tensors to [A,B,2] without new malloc:
[A,2] -> [A,1,2] -> [A,B,2]
[B,2] -> [1,B,2] -> [A,B,2]
Then we compute the area of intersect between box_a and box_b.
Args:
box_a: (tensor) bounding boxes, Shape: [A,4].
box_b: (tensor) bounding boxes, Shape: [B,4].
Return:
(tensor) intersection area, Shape: [A,B].
"""
A = box_a.size(0)
B = box_b.size(0)
max_xy = torch.min(box_a[:, 2:].unsqueeze(1).expand(A, B, 2), box_b[:, 2:].unsqueeze(0).expand(A, B, 2))
min_xy = torch.max(box_a[:, :2].unsqueeze(1).expand(A, B, 2), box_b[:, :2].unsqueeze(0).expand(A, B, 2))
inter = torch.clamp((max_xy - min_xy), min=0)
return inter[:, :, 0] * inter[:, :, 1]
def jaccard(box_a, box_b):
"""Compute the jaccard overlap of two sets of boxes. The jaccard overlap
is simply the intersection over union of two boxes. Here we operate on
ground truth boxes and default boxes.
E.g.:
A ∩ B / A ∪ B = A ∩ B / (area(A) + area(B) - A ∩ B)
Args:
box_a: (tensor) Ground truth bounding boxes, Shape: [num_objects,4]
box_b: (tensor) Prior boxes from priorbox layers, Shape: [num_priors,4]
Return:
jaccard overlap: (tensor) Shape: [box_a.size(0), box_b.size(0)]
"""
inter = intersect(box_a, box_b)
area_a = ((box_a[:, 2] - box_a[:, 0]) * (box_a[:, 3] - box_a[:, 1])).unsqueeze(1).expand_as(inter) # [A,B]
area_b = ((box_b[:, 2] - box_b[:, 0]) * (box_b[:, 3] - box_b[:, 1])).unsqueeze(0).expand_as(inter) # [A,B]
union = area_a + area_b - inter
return inter / union # [A,B]
def matrix_iou(a, b):
"""
return iou of a and b, numpy version for data augenmentation
"""
lt = np.maximum(a[:, np.newaxis, :2], b[:, :2])
rb = np.minimum(a[:, np.newaxis, 2:], b[:, 2:])
area_i = np.prod(rb - lt, axis=2) * (lt < rb).all(axis=2)
area_a = np.prod(a[:, 2:] - a[:, :2], axis=1)
area_b = np.prod(b[:, 2:] - b[:, :2], axis=1)
return area_i / (area_a[:, np.newaxis] + area_b - area_i)
def matrix_iof(a, b):
"""
return iof of a and b, numpy version for data augenmentation
"""
lt = np.maximum(a[:, np.newaxis, :2], b[:, :2])
rb = np.minimum(a[:, np.newaxis, 2:], b[:, 2:])
area_i = np.prod(rb - lt, axis=2) * (lt < rb).all(axis=2)
area_a = np.prod(a[:, 2:] - a[:, :2], axis=1)
return area_i / np.maximum(area_a[:, np.newaxis], 1)
def match(threshold, truths, priors, variances, labels, landms, loc_t, conf_t, landm_t, idx):
"""Match each prior box with the ground truth box of the highest jaccard
overlap, encode the bounding boxes, then return the matched indices
corresponding to both confidence and location preds.
Args:
threshold: (float) The overlap threshold used when matching boxes.
truths: (tensor) Ground truth boxes, Shape: [num_obj, 4].
priors: (tensor) Prior boxes from priorbox layers, Shape: [n_priors,4].
variances: (tensor) Variances corresponding to each prior coord,
Shape: [num_priors, 4].
labels: (tensor) All the class labels for the image, Shape: [num_obj].
landms: (tensor) Ground truth landms, Shape [num_obj, 10].
loc_t: (tensor) Tensor to be filled w/ encoded location targets.
conf_t: (tensor) Tensor to be filled w/ matched indices for conf preds.
landm_t: (tensor) Tensor to be filled w/ encoded landm targets.
idx: (int) current batch index
Return:
The matched indices corresponding to 1)location 2)confidence
3)landm preds.
"""
# jaccard index
overlaps = jaccard(truths, point_form(priors))
# (Bipartite Matching)
# [1,num_objects] best prior for each ground truth
best_prior_overlap, best_prior_idx = overlaps.max(1, keepdim=True)
# ignore hard gt
valid_gt_idx = best_prior_overlap[:, 0] >= 0.2
best_prior_idx_filter = best_prior_idx[valid_gt_idx, :]
if best_prior_idx_filter.shape[0] <= 0:
loc_t[idx] = 0
conf_t[idx] = 0
return
# [1,num_priors] best ground truth for each prior
best_truth_overlap, best_truth_idx = overlaps.max(0, keepdim=True)
best_truth_idx.squeeze_(0)
best_truth_overlap.squeeze_(0)
best_prior_idx.squeeze_(1)
best_prior_idx_filter.squeeze_(1)
best_prior_overlap.squeeze_(1)
best_truth_overlap.index_fill_(0, best_prior_idx_filter, 2) # ensure best prior
# TODO refactor: index best_prior_idx with long tensor
# ensure every gt matches with its prior of max overlap
for j in range(best_prior_idx.size(0)): # 判别此anchor是预测哪一个boxes
best_truth_idx[best_prior_idx[j]] = j
matches = truths[best_truth_idx] # Shape: [num_priors,4] 此处为每一个anchor对应的bbox取出来
conf = labels[best_truth_idx] # Shape: [num_priors] 此处为每一个anchor对应的label取出来
conf[best_truth_overlap < threshold] = 0 # label as background overlap<0.35的全部作为负样本
loc = encode(matches, priors, variances)
matches_landm = landms[best_truth_idx]
landm = encode_landm(matches_landm, priors, variances)
loc_t[idx] = loc # [num_priors,4] encoded offsets to learn
conf_t[idx] = conf # [num_priors] top class label for each prior
landm_t[idx] = landm
def encode(matched, priors, variances):
"""Encode the variances from the priorbox layers into the ground truth boxes
we have matched (based on jaccard overlap) with the prior boxes.
Args:
matched: (tensor) Coords of ground truth for each prior in point-form
Shape: [num_priors, 4].
priors: (tensor) Prior boxes in center-offset form
Shape: [num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
encoded boxes (tensor), Shape: [num_priors, 4]
"""
# dist b/t match center and prior's center
g_cxcy = (matched[:, :2] + matched[:, 2:]) / 2 - priors[:, :2]
# encode variance
g_cxcy /= (variances[0] * priors[:, 2:])
# match wh / prior wh
g_wh = (matched[:, 2:] - matched[:, :2]) / priors[:, 2:]
g_wh = torch.log(g_wh) / variances[1]
# return target for smooth_l1_loss
return torch.cat([g_cxcy, g_wh], 1) # [num_priors,4]
def encode_landm(matched, priors, variances):
"""Encode the variances from the priorbox layers into the ground truth boxes
we have matched (based on jaccard overlap) with the prior boxes.
Args:
matched: (tensor) Coords of ground truth for each prior in point-form
Shape: [num_priors, 10].
priors: (tensor) Prior boxes in center-offset form
Shape: [num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
encoded landm (tensor), Shape: [num_priors, 10]
"""
# dist b/t match center and prior's center
matched = torch.reshape(matched, (matched.size(0), 5, 2))
priors_cx = priors[:, 0].unsqueeze(1).expand(matched.size(0), 5).unsqueeze(2)
priors_cy = priors[:, 1].unsqueeze(1).expand(matched.size(0), 5).unsqueeze(2)
priors_w = priors[:, 2].unsqueeze(1).expand(matched.size(0), 5).unsqueeze(2)
priors_h = priors[:, 3].unsqueeze(1).expand(matched.size(0), 5).unsqueeze(2)
priors = torch.cat([priors_cx, priors_cy, priors_w, priors_h], dim=2)
g_cxcy = matched[:, :, :2] - priors[:, :, :2]
# encode variance
g_cxcy /= (variances[0] * priors[:, :, 2:])
# g_cxcy /= priors[:, :, 2:]
g_cxcy = g_cxcy.reshape(g_cxcy.size(0), -1)
# return target for smooth_l1_loss
return g_cxcy
# Adapted from https://github.com/Hakuyume/chainer-ssd
def decode(loc, priors, variances):
"""Decode locations from predictions using priors to undo
the encoding we did for offset regression at train time.
Args:
loc (tensor): location predictions for loc layers,
Shape: [num_priors,4]
priors (tensor): Prior boxes in center-offset form.
Shape: [num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
decoded bounding box predictions
"""
boxes = torch.cat((priors[:, :2] + loc[:, :2] * variances[0] * priors[:, 2:],
priors[:, 2:] * torch.exp(loc[:, 2:] * variances[1])), 1)
boxes[:, :2] -= boxes[:, 2:] / 2
boxes[:, 2:] += boxes[:, :2]
return boxes
def decode_landm(pre, priors, variances):
"""Decode landm from predictions using priors to undo
the encoding we did for offset regression at train time.
Args:
pre (tensor): landm predictions for loc layers,
Shape: [num_priors,10]
priors (tensor): Prior boxes in center-offset form.
Shape: [num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
decoded landm predictions
"""
tmp = (
priors[:, :2] + pre[:, :2] * variances[0] * priors[:, 2:],
priors[:, :2] + pre[:, 2:4] * variances[0] * priors[:, 2:],
priors[:, :2] + pre[:, 4:6] * variances[0] * priors[:, 2:],
priors[:, :2] + pre[:, 6:8] * variances[0] * priors[:, 2:],
priors[:, :2] + pre[:, 8:10] * variances[0] * priors[:, 2:],
)
landms = torch.cat(tmp, dim=1)
return landms
def batched_decode(b_loc, priors, variances):
"""Decode locations from predictions using priors to undo
the encoding we did for offset regression at train time.
Args:
b_loc (tensor): location predictions for loc layers,
Shape: [num_batches,num_priors,4]
priors (tensor): Prior boxes in center-offset form.
Shape: [1,num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
decoded bounding box predictions
"""
boxes = (
priors[:, :, :2] + b_loc[:, :, :2] * variances[0] * priors[:, :, 2:],
priors[:, :, 2:] * torch.exp(b_loc[:, :, 2:] * variances[1]),
)
boxes = torch.cat(boxes, dim=2)
boxes[:, :, :2] -= boxes[:, :, 2:] / 2
boxes[:, :, 2:] += boxes[:, :, :2]
return boxes
def batched_decode_landm(pre, priors, variances):
"""Decode landm from predictions using priors to undo
the encoding we did for offset regression at train time.
Args:
pre (tensor): landm predictions for loc layers,
Shape: [num_batches,num_priors,10]
priors (tensor): Prior boxes in center-offset form.
Shape: [1,num_priors,4].
variances: (list[float]) Variances of priorboxes
Return:
decoded landm predictions
"""
landms = (
priors[:, :, :2] + pre[:, :, :2] * variances[0] * priors[:, :, 2:],
priors[:, :, :2] + pre[:, :, 2:4] * variances[0] * priors[:, :, 2:],
priors[:, :, :2] + pre[:, :, 4:6] * variances[0] * priors[:, :, 2:],
priors[:, :, :2] + pre[:, :, 6:8] * variances[0] * priors[:, :, 2:],
priors[:, :, :2] + pre[:, :, 8:10] * variances[0] * priors[:, :, 2:],
)
landms = torch.cat(landms, dim=2)
return landms
def log_sum_exp(x):
"""Utility function for computing log_sum_exp while determining
This will be used to determine unaveraged confidence loss across
all examples in a batch.
Args:
x (Variable(tensor)): conf_preds from conf layers
"""
x_max = x.data.max()
return torch.log(torch.sum(torch.exp(x - x_max), 1, keepdim=True)) + x_max
# Original author: Francisco Massa:
# https://github.com/fmassa/object-detection.torch
# Ported to PyTorch by Max deGroot (02/01/2017)
def nms(boxes, scores, overlap=0.5, top_k=200):
"""Apply non-maximum suppression at test time to avoid detecting too many
overlapping bounding boxes for a given object.
Args:
boxes: (tensor) The location preds for the img, Shape: [num_priors,4].
scores: (tensor) The class predscores for the img, Shape:[num_priors].
overlap: (float) The overlap thresh for suppressing unnecessary boxes.
top_k: (int) The Maximum number of box preds to consider.
Return:
The indices of the kept boxes with respect to num_priors.
"""
keep = torch.Tensor(scores.size(0)).fill_(0).long()
if boxes.numel() == 0:
return keep
x1 = boxes[:, 0]
y1 = boxes[:, 1]
x2 = boxes[:, 2]
y2 = boxes[:, 3]
area = torch.mul(x2 - x1, y2 - y1)
v, idx = scores.sort(0) # sort in ascending order
# I = I[v >= 0.01]
idx = idx[-top_k:] # indices of the top-k largest vals
xx1 = boxes.new()
yy1 = boxes.new()
xx2 = boxes.new()
yy2 = boxes.new()
w = boxes.new()
h = boxes.new()
# keep = torch.Tensor()
count = 0
while idx.numel() > 0:
i = idx[-1] # index of current largest val
# keep.append(i)
keep[count] = i
count += 1
if idx.size(0) == 1:
break
idx = idx[:-1] # remove kept element from view
# load bboxes of next highest vals
torch.index_select(x1, 0, idx, out=xx1)
torch.index_select(y1, 0, idx, out=yy1)
torch.index_select(x2, 0, idx, out=xx2)
torch.index_select(y2, 0, idx, out=yy2)
# store element-wise max with next highest score
xx1 = torch.clamp(xx1, min=x1[i])
yy1 = torch.clamp(yy1, min=y1[i])
xx2 = torch.clamp(xx2, max=x2[i])
yy2 = torch.clamp(yy2, max=y2[i])
w.resize_as_(xx2)
h.resize_as_(yy2)
w = xx2 - xx1
h = yy2 - yy1
# check sizes of xx1 and xx2.. after each iteration
w = torch.clamp(w, min=0.0)
h = torch.clamp(h, min=0.0)
inter = w * h
# IoU = i / (area(a) + area(b) - i)
rem_areas = torch.index_select(area, 0, idx) # load remaining areas)
union = (rem_areas - inter) + area[i]
IoU = inter / union # store result in iou
# keep only elements with an IoU <= overlap
idx = idx[IoU.le(overlap)]
return keep, count