mirror of
https://github.com/jiaxiaojunQAQ/OmniSafeBench-MM.git
synced 2026-08-22 03:07:15 +02:00
refactor: consolidate duplicate code and fix deadlock bug
- Extract duplicate attack/defense config merging into _merge_component_configs() - Extract duplicate lazy loading logic into _get_component() - Move content policy detection to BaseModel base class - Fix BatchSaveManager deadlock by splitting flush logic - Add TypeError to ValueError conversion for consistent config errors - Move _determine_load_model() to BaseComponent (explicit field only)
This commit is contained in:
+23
-39
@@ -47,7 +47,11 @@ class BaseComponent(ABC):
|
||||
cfg_dict = config
|
||||
allowed_fields = {f.name for f in fields(self.CONFIG_CLASS)}
|
||||
filtered = {k: v for k, v in cfg_dict.items() if k in allowed_fields}
|
||||
self.cfg = self.CONFIG_CLASS(**filtered)
|
||||
try:
|
||||
self.cfg = self.CONFIG_CLASS(**filtered)
|
||||
except TypeError as e:
|
||||
# Convert TypeError to ValueError for consistency
|
||||
raise ValueError(f"Invalid configuration: {e}") from e
|
||||
else:
|
||||
# No configuration class or not a dataclass, use dictionary directly
|
||||
self.cfg = config
|
||||
@@ -77,6 +81,24 @@ class BaseComponent(ABC):
|
||||
pass
|
||||
# Subclasses can add more validation logic
|
||||
|
||||
def _determine_load_model(self) -> bool:
|
||||
"""Determine if local model needs to be loaded.
|
||||
|
||||
Checks the load_model field in configuration (dict or dataclass).
|
||||
Returns False if load_model is not explicitly set.
|
||||
"""
|
||||
if hasattr(self, "cfg"):
|
||||
config_obj = self.cfg
|
||||
else:
|
||||
config_obj = self.config
|
||||
|
||||
if isinstance(config_obj, dict):
|
||||
return config_obj.get("load_model", False)
|
||||
elif hasattr(config_obj, "load_model"):
|
||||
return getattr(config_obj, "load_model", False)
|
||||
|
||||
return False
|
||||
|
||||
|
||||
class BaseAttack(BaseComponent, ABC):
|
||||
"""Attack method base class (enhanced version)"""
|
||||
@@ -109,25 +131,6 @@ class BaseAttack(BaseComponent, ABC):
|
||||
# Determine if local model needs to be loaded
|
||||
self.load_model = self._determine_load_model()
|
||||
|
||||
def _determine_load_model(self) -> bool:
|
||||
"""Determine if local model needs to be loaded
|
||||
|
||||
Only decide based on the load_model field in configuration
|
||||
"""
|
||||
# Check if load_model field exists in configuration
|
||||
if hasattr(self, "cfg"):
|
||||
config_obj = self.cfg
|
||||
else:
|
||||
config_obj = self.config
|
||||
|
||||
# Only check load_model configuration item
|
||||
if isinstance(config_obj, dict):
|
||||
return config_obj.get("load_model", False)
|
||||
elif hasattr(config_obj, "load_model"):
|
||||
return getattr(config_obj, "load_model", False)
|
||||
|
||||
return False
|
||||
|
||||
@abstractmethod
|
||||
def generate_test_case(
|
||||
self, original_prompt: str, image_path: str, case_id: str, **kwargs
|
||||
@@ -259,25 +262,6 @@ class BaseDefense(BaseComponent, ABC):
|
||||
# Determine if local model needs to be loaded
|
||||
self.load_model = self._determine_load_model()
|
||||
|
||||
def _determine_load_model(self) -> bool:
|
||||
"""Determine if local model needs to be loaded
|
||||
|
||||
Only based on the load_model field in configuration
|
||||
"""
|
||||
# Check if load_model field exists in configuration
|
||||
if hasattr(self, "cfg"):
|
||||
config_obj = self.cfg
|
||||
else:
|
||||
config_obj = self.config
|
||||
|
||||
# Only check load_model configuration item
|
||||
if isinstance(config_obj, dict):
|
||||
return config_obj.get("load_model", False)
|
||||
elif hasattr(config_obj, "load_model"):
|
||||
return getattr(config_obj, "load_model", False)
|
||||
|
||||
return False
|
||||
|
||||
@abstractmethod
|
||||
def apply_defense(self, test_case: TestCase, **kwargs) -> TestCase:
|
||||
"""
|
||||
|
||||
Reference in New Issue
Block a user