diff --git a/AILab_RMBG.py b/AILab_RMBG.py index f8df678..d82ad28 100644 --- a/AILab_RMBG.py +++ b/AILab_RMBG.py @@ -1,4 +1,4 @@ -# ComfyUI-RMBG v2.1.0 +# ComfyUI-RMBG v2.1.1 # 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. @@ -141,7 +141,7 @@ class BaseModelLoader: if self.model is not None: self.model.cpu() del self.model - # 实际有用的内存清理 + import gc gc.collect() if torch.cuda.is_available(): @@ -158,11 +158,46 @@ class RMBGModel(BaseModelLoader): self.clear_model() cache_dir = self.get_cache_dir(model_name) - self.model = AutoModelForImageSegmentation.from_pretrained( - cache_dir, - trust_remote_code=True, - local_files_only=True - ) + try: + self.model = AutoModelForImageSegmentation.from_pretrained( + cache_dir, + trust_remote_code=True, + local_files_only=True + ) + except Exception as e: + if "'Config' object has no attribute 'get_text_config'" in str(e): + print("[RMBG WARNING] Detected newer transformers version, attempting 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") + module_name = f"custom_birefnet_model_{hash(birefnet_path)}" + spec = importlib.util.spec_from_file_location(module_name, birefnet_path) + birefnet_module = importlib.util.module_from_spec(spec) + sys.modules[module_name] = birefnet_module + spec.loader.exec_module(birefnet_module) + + for attr_name in dir(birefnet_module): + attr = getattr(birefnet_module, attr_name) + if isinstance(attr, type) and issubclass(attr, PreTrainedModel) and attr != PreTrainedModel: + config_class = getattr(birefnet_module, "BiRefNet_config", None) + if config_class: + model_config = config_class() + self.model = attr(model_config) + self.model.load_state_dict(torch.load(os.path.join(cache_dir, "model.safetensors"))) + 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)}\nConsider downgrading transformers to version 4.48.3: pip install transformers==4.48.3") + else: + raise e self.model.eval() for param in self.model.parameters(): @@ -198,7 +233,26 @@ class RMBGModel(BaseModelLoader): input_batch = torch.cat(input_tensors, dim=0).to(device) with torch.no_grad(): - results = self.model(input_batch)[-1].sigmoid().cpu() + 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