Load a LoRA adapter with enhanced caching
(self, base_model: PreTrainedModel, adapter_id: str)
| 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, |
| 1171 | adapter_name = adapter_name, |
| 1172 | ) |
| 1173 | |
| 1174 | if model not in self.loaded_adapters: |
| 1175 | self.loaded_adapters[model] = [] |
| 1176 | if adapter_id not in self.loaded_adapters[model]: |
| 1177 | self.loaded_adapters[model].append(adapter_id) |
| 1178 | |
| 1179 | return model |
| 1180 | |
| 1181 | except Exception as e: |
| 1182 | error_msg = f"Failed to load adapter {adapter_id}: {str(e)}" |
| 1183 | logger.error(error_msg) |
| 1184 | raise RuntimeError(error_msg) from e |
| 1185 | |
| 1186 | return self.cache_manager.get_or_load_adapter(model_key, adapter_id, _load_adapter) |
| 1187 | |
| 1188 | def set_active_adapter(self, model: PeftModel, adapter_id: str = None) -> bool: |
| 1189 | """Set a specific adapter as active with error handling""" |
no test coverage detected