MCPcopy Create free account
hub / github.com/BIT-DataLab/Edit-Banana / RMBGModel

Class RMBGModel

modules/icon_picture_processor.py:37–226  ·  view source on GitHub ↗

RMBG-2.0 background-removal model (ONNX Runtime, CUDA if available). Example: model = RMBGModel(model_path) model.load() rgba_image = model.remove_background(pil_image)

Source from the content-addressed store, hash-verified

35
36# ======================== RMBG-2.0 model wrapper ========================
37class RMBGModel(ModelWrapper):
38 """
39 RMBG-2.0 background-removal model (ONNX Runtime, CUDA if available).
40
41 Example:
42 model = RMBGModel(model_path)
43 model.load()
44 rgba_image = model.remove_background(pil_image)
45 """
46
47 INPUT_SIZE = (1024, 1024)
48
49 def __init__(self, model_path: str = None):
50 super().__init__()
51 self.model_path = model_path or self._get_default_path()
52 self._session = None
53 self._input_name = None
54 self._output_name = None
55
56 def _get_default_path(self) -> str:
57 """Default model path under models/rmbg/."""
58 return os.path.join(
59 os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
60 "models", "rmbg", "model.onnx"
61 )
62
63 def load(self):
64 """Load RMBG-2.0 ONNX model; fallback to CPU if CUDA fails."""
65 if self._is_loaded:
66 return
67
68 if not ONNX_AVAILABLE:
69 print("[RMBGModel] Warning: onnxruntime not available, using fallback mode")
70 self._is_loaded = True
71 return
72
73 if not os.path.exists(self.model_path):
74 print(f"[RMBGModel] Warning: Model file not found at {self.model_path}, using fallback mode")
75 self._is_loaded = True
76 return
77
78 # ONNX Runtime options
79 session_options = ort.SessionOptions()
80 session_options.log_severity_level = 3 # ERROR only
81 session_options.enable_profiling = False
82
83 # Available providers
84 available_providers = ort.get_available_providers()
85
86 # Try CUDA then CPU
87 providers_to_try = [
88 (['CUDAExecutionProvider', 'CPUExecutionProvider'], "CUDA+CPU"),
89 (['CPUExecutionProvider'], "CPU only"),
90 ]
91
92 for providers, name in providers_to_try:
93 # Filter valid providers
94 valid_providers = [p for p in providers if p in available_providers]

Callers 1

load_rmbg_modelMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected