v1.0.0
@@ -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 }}
|
||||
@@ -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)"
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
[中文](README-CN.md)|[English](README.md)
|
||||
|
||||
# 人像处理等相关的 ComfyUI 节点
|
||||
|
||||
目前包含以下节点:
|
||||
- 图片人脸对齐(正脸);
|
||||
- 人脸检测裁剪, 可选是否对齐, 可调裁剪区域大小, 角度;
|
||||
- 各种证件照一键生成;
|
||||
- 美化照片, 包括亮度, 饱和度, 锐化, 磨皮等.
|
||||
|
||||
示例:
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||

|
||||
|
||||

|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
|
||||
## 📣 更新
|
||||
|
||||
[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)
|
||||
@@ -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:
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||

|
||||
|
||||

|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
## 📣 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)
|
||||
@@ -0,0 +1,3 @@
|
||||
from .ptnodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
@@ -0,0 +1,3 @@
|
||||
from beauty.base_adjust import *
|
||||
from beauty.grind_skin import *
|
||||
from beauty.whitening import *
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
After Width: | Height: | Size: 103 KiB |
@@ -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)
|
||||
|
||||
|
After Width: | Height: | Size: 174 KiB |
|
After Width: | Height: | Size: 270 KiB |
|
After Width: | Height: | Size: 190 KiB |
|
After Width: | Height: | Size: 406 KiB |
|
After Width: | Height: | Size: 222 KiB |
|
After Width: | Height: | Size: 398 KiB |
@@ -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
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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 = ""
|
||||
@@ -0,0 +1,2 @@
|
||||
huggingface_hub
|
||||
opencv-python
|
||||
@@ -0,0 +1 @@
|
||||
from retinaface.retinaface import RetinaFace
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||