diff --git a/.github/workflows/publish_action.yml b/.github/workflows/publish_action.yml new file mode 100644 index 0000000..ef9d8d9 --- /dev/null +++ b/.github/workflows/publish_action.yml @@ -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 }} diff --git a/AILab_RMBG.py b/AILab_RMBG.py new file mode 100644 index 0000000..38b12d9 --- /dev/null +++ b/AILab_RMBG.py @@ -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)" +} \ No newline at end of file diff --git a/README-CN.md b/README-CN.md new file mode 100644 index 0000000..c7461f1 --- /dev/null +++ b/README-CN.md @@ -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) diff --git a/README.md b/README.md index 72eb711..472240a 100644 --- a/README.md +++ b/README.md @@ -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) \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..94be845 --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .ptnodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/beauty/__init__.py b/beauty/__init__.py new file mode 100644 index 0000000..ecd79f2 --- /dev/null +++ b/beauty/__init__.py @@ -0,0 +1,3 @@ +from beauty.base_adjust import * +from beauty.grind_skin import * +from beauty.whitening import * \ No newline at end of file diff --git a/beauty/base_adjust.py b/beauty/base_adjust.py new file mode 100644 index 0000000..4eab345 --- /dev/null +++ b/beauty/base_adjust.py @@ -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 diff --git a/beauty/grind_skin.py b/beauty/grind_skin.py new file mode 100644 index 0000000..20dc16a --- /dev/null +++ b/beauty/grind_skin.py @@ -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 + diff --git a/beauty/lut/lut_origin.png b/beauty/lut/lut_origin.png new file mode 100644 index 0000000..743fc12 Binary files /dev/null and b/beauty/lut/lut_origin.png differ diff --git a/beauty/whitening.py b/beauty/whitening.py new file mode 100644 index 0000000..9ca8503 --- /dev/null +++ b/beauty/whitening.py @@ -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) + diff --git a/images/2025-04-11_07-06-36.png b/images/2025-04-11_07-06-36.png new file mode 100644 index 0000000..9088ce3 Binary files /dev/null and b/images/2025-04-11_07-06-36.png differ diff --git a/images/2025-04-11_07-08-46.png b/images/2025-04-11_07-08-46.png new file mode 100644 index 0000000..5572c1f Binary files /dev/null and b/images/2025-04-11_07-08-46.png differ diff --git a/images/2025-04-11_07-10-24.png b/images/2025-04-11_07-10-24.png new file mode 100644 index 0000000..68baf51 Binary files /dev/null and b/images/2025-04-11_07-10-24.png differ diff --git a/images/2025-04-11_09-05-41.png b/images/2025-04-11_09-05-41.png new file mode 100644 index 0000000..7d8e57e Binary files /dev/null and b/images/2025-04-11_09-05-41.png differ diff --git a/images/2025-04-11_09-27-16.png b/images/2025-04-11_09-27-16.png new file mode 100644 index 0000000..d7c25ea Binary files /dev/null and b/images/2025-04-11_09-27-16.png differ diff --git a/images/2025-04-11_09-48-23.png b/images/2025-04-11_09-48-23.png new file mode 100644 index 0000000..826c720 Binary files /dev/null and b/images/2025-04-11_09-48-23.png differ diff --git a/layout_calculator.py b/layout_calculator.py new file mode 100644 index 0000000..d034679 --- /dev/null +++ b/layout_calculator.py @@ -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 diff --git a/ptnodes.py b/ptnodes.py new file mode 100644 index 0000000..3896a66 --- /dev/null +++ b/ptnodes.py @@ -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", +} \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..8bdd6bf --- /dev/null +++ b/pyproject.toml @@ -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 = "" diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..10bee69 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,2 @@ +huggingface_hub +opencv-python diff --git a/retinaface/__init__.py b/retinaface/__init__.py new file mode 100644 index 0000000..5ba987e --- /dev/null +++ b/retinaface/__init__.py @@ -0,0 +1 @@ +from retinaface.retinaface import RetinaFace \ No newline at end of file diff --git a/retinaface/retinaface.py b/retinaface/retinaface.py new file mode 100644 index 0000000..29a41c4 --- /dev/null +++ b/retinaface/retinaface.py @@ -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 diff --git a/retinaface/retinaface_net.py b/retinaface/retinaface_net.py new file mode 100644 index 0000000..c52535e --- /dev/null +++ b/retinaface/retinaface_net.py @@ -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 diff --git a/retinaface/retinaface_utils.py b/retinaface/retinaface_utils.py new file mode 100644 index 0000000..f19e320 --- /dev/null +++ b/retinaface/retinaface_utils.py @@ -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