LoRA manager with enhanced error handling and caching
| 1111 | return self.cache_manager.get_or_load_model(model_id, _load_model) |
| 1112 | |
| 1113 | class LoRAManager: |
| 1114 | """LoRA manager with enhanced error handling and caching""" |
| 1115 | |
| 1116 | def __init__(self, cache_manager: CacheManager): |
| 1117 | self.cache_manager = cache_manager |
| 1118 | self.loaded_adapters = {} |
| 1119 | self.adapter_names = {} # Maps adapter_id to valid adapter name |
| 1120 | |
| 1121 | def _get_adapter_name(self, adapter_id: str) -> str: |
| 1122 | """Create a valid adapter name from adapter_id.""" |
| 1123 | if adapter_id in self.adapter_names: |
| 1124 | return self.adapter_names[adapter_id] |
| 1125 | |
| 1126 | name = adapter_id.replace('.', '_').replace('-', '_') |
| 1127 | name = ''.join(c if c.isalnum() or c == '_' else '' for c in name) |
| 1128 | if name[0].isdigit(): |
| 1129 | name = f"adapter_{name}" |
| 1130 | |
| 1131 | self.adapter_names[adapter_id] = name |
| 1132 | return name |
| 1133 | |
| 1134 | def validate_adapter(self, adapter_id: str) -> bool: |
| 1135 | """Validate if adapter exists and is compatible""" |
| 1136 | try: |
| 1137 | config = PeftConfig.from_pretrained( |
| 1138 | adapter_id, |
| 1139 | trust_remote_code=True, |
| 1140 | token=os.getenv("HF_TOKEN") |
| 1141 | ) |
| 1142 | return True |
| 1143 | except Exception as e: |
| 1144 | logger.error(f"Error validating adapter {adapter_id}: {str(e)}") |
| 1145 | return False |
| 1146 | |
| 1147 | def load_adapter(self, base_model: PreTrainedModel, adapter_id: str) -> PreTrainedModel: |
| 1148 | """Load a LoRA adapter with enhanced caching""" |
| 1149 | model_key = base_model.config._name_or_path |
| 1150 | |
| 1151 | def _load_adapter(): |
| 1152 | logger.info(f"Loading LoRA adapter: {adapter_id}") |
| 1153 | |
| 1154 | if not self.validate_adapter(adapter_id): |
| 1155 | error_msg = f"Adapter {adapter_id} not found or is not compatible" |
| 1156 | logger.error(error_msg) |
| 1157 | raise ValueError(error_msg) |
| 1158 | |
| 1159 | try: |
| 1160 | adapter_name = self._get_adapter_name(adapter_id) |
| 1161 | |
| 1162 | config = PeftConfig.from_pretrained( |
| 1163 | adapter_id, |
| 1164 | trust_remote_code=True, |
| 1165 | token=os.getenv("HF_TOKEN") |
| 1166 | ) |
| 1167 | |
| 1168 | model = base_model |
| 1169 | model.add_adapter( |
| 1170 | config, |