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)
| 35 | |
| 36 | # ======================== RMBG-2.0 model wrapper ======================== |
| 37 | class 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] |