Load RMBG-2.0 ONNX model; fallback to CPU if CUDA fails.
(self)
| 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] |
| 95 | if not valid_providers: |
| 96 | continue |
| 97 | |
| 98 | try: |
| 99 | print(f"[RMBGModel] Trying to load with {name} ({valid_providers})...") |
| 100 | self._session = ort.InferenceSession( |
| 101 | self.model_path, |
| 102 | providers=valid_providers, |
| 103 | sess_options=session_options |
| 104 | ) |
| 105 | |
| 106 | self._input_name = self._session.get_inputs()[0].name |
| 107 | self._output_name = self._session.get_outputs()[0].name |
| 108 | self._providers = valid_providers |
| 109 | |
| 110 | self._is_loaded = True |
| 111 | print(f"[RMBGModel] Model loaded successfully with {name}") |
| 112 | return |
| 113 | |
| 114 | except Exception as e: |
| 115 | print(f"[RMBGModel] Failed to load with {name}: {e}") |
| 116 | # Try next config |
| 117 | continue |
| 118 | |
| 119 | # All attempts failed, use fallback |
| 120 | print("[RMBGModel] Warning: All loading attempts failed, using fallback mode (no background removal)") |
no outgoing calls
no test coverage detected