From aecd554128177aa08070f5ba4d73e4334cb0ee2f Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Wed, 20 May 2026 08:56:16 -0600 Subject: [PATCH] Add sapiens2 matting as a mask generator. Begin transition to model paths and model folders. --- .../captioner/BaseCaptioner.py | 1 + toolkit/config_modules.py | 2 +- toolkit/control_generator.py | 14 +- toolkit/dataloader_mixins.py | 2 +- toolkit/models/loaders/comfy.py | 15 -- toolkit/models/loaders/umt5.py | 16 +- toolkit/models/sapiens2.py | 190 +++++++++++++++++- toolkit/paths.py | 3 - 8 files changed, 209 insertions(+), 34 deletions(-) delete mode 100644 toolkit/models/loaders/comfy.py diff --git a/extensions_built_in/captioner/BaseCaptioner.py b/extensions_built_in/captioner/BaseCaptioner.py index bd946292..96a13535 100644 --- a/extensions_built_in/captioner/BaseCaptioner.py +++ b/extensions_built_in/captioner/BaseCaptioner.py @@ -149,6 +149,7 @@ class BaseCaptioner(BaseExtensionProcess): def find_files(self): # recursivly find all the files in the path_to_caption with the specified extensions and save the paths to self.file_paths for root, dirs, files in os.walk(self.caption_config.path_to_caption): + dirs[:] = [d for d in dirs if d != "_controls"] for file in files: if any( file.lower().endswith(f".{ext}") and not file.startswith(".") diff --git a/toolkit/config_modules.py b/toolkit/config_modules.py index b2261c5b..cf426524 100644 --- a/toolkit/config_modules.py +++ b/toolkit/config_modules.py @@ -859,7 +859,7 @@ class SliderConfig: self.targets.append(target) print(f"Built {len(self.targets)} slider targets (with permutations)") -ControlTypes = Literal['depth', 'line', 'pose', 'inpaint', 'mask'] +ControlTypes = Literal['depth', 'line', 'pose', 'inpaint', 'mask', 'sapiens2_mask'] class DatasetConfig: """ diff --git a/toolkit/control_generator.py b/toolkit/control_generator.py index 39acfa01..cf7ddaf2 100644 --- a/toolkit/control_generator.py +++ b/toolkit/control_generator.py @@ -148,7 +148,7 @@ class ControlGenerator: img = img.convert('RGB') img.save(save_path) return save_path - elif control_type == 'inpaint' or control_type == 'mask': + elif control_type in ['inpaint', 'mask']: self.debug_print("Generating inpaint/mask control") img = image.copy() if self.control_bg_remover is None: @@ -188,6 +188,18 @@ class ControlGenerator: img = img.convert('RGB') img.save(save_path) return save_path + elif control_type in ['sapiens2_mask']: + self.debug_print("Generating sapiens2_mask control") + if self.control_bg_remover is None: + from toolkit.models.sapiens2 import Sapiens2Matting + self.control_bg_remover = Sapiens2Matting.from_pretrained( + device=device, + dtype=torch.float16 + ) + img = image.copy() + img = self.control_bg_remover(img) + img.save(save_path) + return save_path else: raise Exception(f"Error: unknown control type {control_type}") diff --git a/toolkit/dataloader_mixins.py b/toolkit/dataloader_mixins.py index ffd2f862..02d8ebab 100644 --- a/toolkit/dataloader_mixins.py +++ b/toolkit/dataloader_mixins.py @@ -2297,7 +2297,7 @@ class ControlCachingMixin: if control_type == 'inpaint': file_item.inpaint_path = control_path file_item.has_inpaint_image = True - elif control_type == 'mask': + elif control_type == 'mask' or control_type == 'sapiens2_mask': file_item.mask_path = control_path file_item.has_mask_image = True else: diff --git a/toolkit/models/loaders/comfy.py b/toolkit/models/loaders/comfy.py deleted file mode 100644 index 38a32046..00000000 --- a/toolkit/models/loaders/comfy.py +++ /dev/null @@ -1,15 +0,0 @@ -import os -from typing import List -from toolkit.paths import COMFY_MODELS_PATH - - -def get_comfy_path(comfy_files: List[str]) -> str: - """ - Get the path to the first existing file in the COMFY_MODELS_PATH. - """ - if COMFY_MODELS_PATH is not None and comfy_files is not None and len(comfy_files) > 0: - for file in comfy_files: - file_path = os.path.join(COMFY_MODELS_PATH, file) - if os.path.exists(file_path): - return file_path - return None \ No newline at end of file diff --git a/toolkit/models/loaders/umt5.py b/toolkit/models/loaders/umt5.py index b3613923..f53db3cb 100644 --- a/toolkit/models/loaders/umt5.py +++ b/toolkit/models/loaders/umt5.py @@ -1,7 +1,6 @@ from typing import List import torch from transformers import T5Tokenizer, UMT5EncoderModel -from toolkit.models.loaders.comfy import get_comfy_path class PatchedT5Tokenizer(T5Tokenizer): def __init__( @@ -40,15 +39,8 @@ def get_umt5_encoder( Load the UMT5 encoder model from the specified path. """ tokenizer = PatchedT5Tokenizer.from_pretrained(model_path, subfolder=tokenizer_subfolder) - comfy_path = get_comfy_path(comfy_files) - comfy_path = None - if comfy_path is not None: - text_encoder = UMT5EncoderModel.from_single_file( - comfy_path, torch_dtype=torch_dtype - ) - else: - print(f"Using {model_path} for UMT5 encoder.") - text_encoder = UMT5EncoderModel.from_pretrained( - model_path, subfolder=encoder_subfolder, torch_dtype=torch_dtype - ) + print(f"Using {model_path} for UMT5 encoder.") + text_encoder = UMT5EncoderModel.from_pretrained( + model_path, subfolder=encoder_subfolder, torch_dtype=torch_dtype + ) return tokenizer, text_encoder diff --git a/toolkit/models/sapiens2.py b/toolkit/models/sapiens2.py index e1b9d51b..9f4fb28a 100644 --- a/toolkit/models/sapiens2.py +++ b/toolkit/models/sapiens2.py @@ -7,7 +7,7 @@ # https://raw.githubusercontent.com/facebookresearch/sapiens2/refs/heads/main/sapiens/backbones/standalone/sapiens2.py import math -from typing import Any, Dict, List, Literal, Optional, Sequence, Tuple, Union +from typing import Any, List, Literal, Optional, Sequence, Tuple, Union import torch import torch.nn as nn @@ -15,6 +15,8 @@ import torch.nn.functional as F from torch import Tensor from torch.nn.init import trunc_normal_ from torch.utils.checkpoint import checkpoint +from toolkit.paths import MODELS_PATH +import os # ---------------------------------------------------------------------------- @@ -942,3 +944,189 @@ def imagenet_normalize(tensors_0_1: torch.Tensor) -> torch.Tensor: _IMAGENET_STD, dtype=tensors_0_1.dtype, device=tensors_0_1.device ).view(1, 3, 1, 1) return (tensors_0_1 - mean) / std + + +# ---------------------------------------------------------------------------- +class MattingHead(nn.Module): + """Matting decode head from + https://github.com/facebookresearch/sapiens2/blob/main/sapiens/dense/src/models/heads/matting_head.py + + Predicts a 4-channel output: pre-multiplied foreground RGB (channels 0-2) + and soft alpha matte (channel 3), all in [0, 1] after sigmoid. + """ + + def __init__( + self, + in_channels: int = 1536, + upsample_channels: Sequence[int] = (768, 512, 256, 128), + conv_out_channels: Optional[Sequence[int]] = (64, 32, 16), + conv_kernel_sizes: Optional[Sequence[int]] = (3, 3, 3), + out_channels: int = 4, + ): + super().__init__() + self.in_channels = in_channels + + self.input_conv = nn.Sequential( + nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1), + nn.InstanceNorm2d(in_channels), + nn.SiLU(inplace=True), + ) + + up_blocks = [] + cur_ch = in_channels + for out_ch in upsample_channels: + up_blocks.append( + nn.Sequential( + nn.Conv2d(cur_ch, out_ch * 4, kernel_size=3, padding=1), + nn.PixelShuffle(2), + nn.InstanceNorm2d(out_ch), + nn.SiLU(inplace=True), + ) + ) + cur_ch = out_ch + self.upsample_blocks = nn.Sequential(*up_blocks) + + conv_layers = [] + if conv_out_channels and conv_kernel_sizes: + for out_ch, k in zip(conv_out_channels, conv_kernel_sizes): + conv_layers.extend( + [ + nn.Conv2d(cur_ch, out_ch, k, padding=(k - 1) // 2), + nn.InstanceNorm2d(out_ch), + nn.SiLU(inplace=True), + ] + ) + cur_ch = out_ch + self.conv_layers = nn.Sequential(*conv_layers) + + self.conv_matting = nn.Conv2d(cur_ch, out_channels, kernel_size=1) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.input_conv(x) + x = self.upsample_blocks(x) + x = self.conv_layers(x) + return self.conv_matting(x).sigmoid() + + +# ---------------------------------------------------------------------------- +class Sapiens2Matting(nn.Module): + """Sapiens2 backbone + MattingHead for human image matting. + + Reference: https://github.com/facebookresearch/sapiens2/blob/main/docs/MATTING.md + """ + + _ARCH_TO_EMBED_DIM = { + "sapiens2_0.1b": 768, + "sapiens2_0.4b": 1024, + "sapiens2_0.8b": 1280, + "sapiens2_1b": 1536, + "sapiens2_5b": 2432, + } + + def __init__( + self, + arch: str = "sapiens2_1b", + img_size: Tuple[int, int] = (1024, 768), + patch_size: int = 16, + ): + super().__init__() + arch = arch.lower() + if arch not in self._ARCH_TO_EMBED_DIM: + raise ValueError(f"Unsupported arch {arch}") + + self.arch = arch + self.img_size = to_2tuple(img_size) + self.patch_size = patch_size + + self.backbone = Sapiens2( + arch=arch, + img_size=img_size, + patch_size=patch_size, + final_norm=True, + use_tokenizer=False, + with_cls_token=True, + out_type="featmap", + ) + self.decode_head = MattingHead( + in_channels=self._ARCH_TO_EMBED_DIM[arch], + upsample_channels=(768, 512, 256, 128), + conv_out_channels=(64, 32, 16), + conv_kernel_sizes=(3, 3, 3), + out_channels=4, + ) + + @classmethod + def from_pretrained( + cls, + repo_id: str = "facebook/sapiens2-matting-1b", + filename: str = "sapiens2_1b_matting.safetensors", + arch: str = "sapiens2_1b", + img_size: Tuple[int, int] = (1024, 768), + patch_size: int = 16, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + ) -> "Sapiens2Matting": + import huggingface_hub + from safetensors.torch import load_file + + safetensors_path = os.path.join(MODELS_PATH, "sapiens2", filename) + if not os.path.exists(safetensors_path): + print(f"Downloading pretrained weights from HuggingFace Hub: {repo_id}/{filename}...") + os.makedirs(os.path.dirname(safetensors_path), exist_ok=True) + safetensors_path = huggingface_hub.hf_hub_download( + repo_id=repo_id, + filename=filename, + local_dir=os.path.join(MODELS_PATH, "sapiens2"), + ) + model = cls(arch=arch, img_size=img_size, patch_size=patch_size) + state_dict = load_file(safetensors_path) + model.load_state_dict(state_dict) + model.eval() + if device is not None or dtype is not None: + model.to(device=device, dtype=dtype) + return model + + @property + def device(self): + return next(self.parameters()).device + + @property + def dtype(self): + return next(self.parameters()).dtype + + @torch.no_grad() + def forward(self, image, max_res: int = 1024): + """Take a PIL image and return a PIL alpha-matte mask in RGB mode at + the original input size. The image is run through the model at its + native aspect ratio, snapped to a multiple of patch_size and capped + at max_res*max_res pixels.""" + from torchvision import transforms + + p = self.patch_size + w, h = image.size + target_h, target_w = h, w + if target_h * target_w > max_res * max_res: + scale = math.sqrt((max_res * max_res) / (target_h * target_w)) + target_h = int(target_h * scale) + target_w = int(target_w * scale) + target_h = max(p, (target_h // p) * p) + target_w = max(p, (target_w // p) * p) + + transform_image = transforms.Compose( + [ + transforms.Resize((target_h, target_w)), + transforms.ToTensor(), + transforms.Normalize(_IMAGENET_MEAN, _IMAGENET_STD), + ] + ) + input_images = ( + transform_image(image).unsqueeze(0).to(self.device, dtype=self.dtype) + ) + + feat = self.backbone(input_images)[0] + out = self.decode_head(feat) # (1, 4, H, W) in [0, 1] + alpha = out[0, 3].float().cpu() + + mask = transforms.ToPILImage()(alpha) + mask = mask.resize(image.size).convert("RGB") + return mask diff --git a/toolkit/paths.py b/toolkit/paths.py index edd36ce1..9f6a8cfb 100644 --- a/toolkit/paths.py +++ b/toolkit/paths.py @@ -5,10 +5,7 @@ CONFIG_ROOT = os.path.join(TOOLKIT_ROOT, 'config') KEYMAPS_ROOT = os.path.join(TOOLKIT_ROOT, "toolkit", "keymaps") ORIG_CONFIGS_ROOT = os.path.join(TOOLKIT_ROOT, "toolkit", "orig_configs") DIFFUSERS_CONFIGS_ROOT = os.path.join(TOOLKIT_ROOT, "toolkit", "diffusers_configs") -COMFY_PATH = os.getenv("COMFY_PATH", None) COMFY_MODELS_PATH = None -if COMFY_PATH: - COMFY_MODELS_PATH = os.path.join(COMFY_PATH, "models") # check if ENV variable is set if 'MODELS_PATH' in os.environ: