| import peft |
| import torch |
| import logging |
|
|
| from contextlib import contextmanager |
|
|
|
|
| @contextmanager |
| def use_lora_adapter( |
| model, |
| lora_paths: str | list[str], |
| adapter_names: str | list[str], |
| merge_weights: float | list[float], |
| logger: logging.Logger, |
| ): |
| |
| if isinstance(lora_paths, str): |
| lora_paths = [lora_paths] |
| if isinstance(adapter_names, str): |
| adapter_names = [adapter_names] |
| if isinstance(merge_weights, float): |
| merge_weights = [merge_weights] |
|
|
| transformer = model.transformer |
| while isinstance(transformer, peft.PeftModel): |
| transformer = transformer.base_model.model |
| |
| if isinstance(model.transformer, peft.PeftModel): |
| logger.info("[LoRA] Unloading existing PeftModel wrapper.") |
| try: |
| model.transformer = model.transformer.unload() |
| except AttributeError: |
| model.transformer = model.transformer.base_model.model |
| |
| if hasattr(model.transformer, "peft_config"): |
| del model.transformer.peft_config |
| logger.info(f"[LoRA] Successfully removed residual peft_config from Base Model.") |
| |
| transformer = model.transformer |
|
|
| try: |
| if torch.cuda.is_available(): |
| mem_before = torch.cuda.memory_allocated(model.device) / 1024**2 |
| logger.info(f"[LoRA] GPU memory before load: {mem_before:.2f} MB") |
|
|
| configs = [peft.LoraConfig.from_pretrained(p) for p in lora_paths] |
| for i in range(1, len(configs)): |
| assert configs[i] == configs[0], "All LoRA configs must be identical to merge." |
| lora_config = configs[0] |
| |
| transformer = peft.get_peft_model(transformer, lora_config, adapter_name=adapter_names[0]) |
| logger.info(f"After get_peft_model: {type(transformer)=}") |
|
|
| for path, name in zip(lora_paths, adapter_names): |
| logger.info(f"[LoRA] Loading adapter: {name=}, {path=}") |
| transformer.load_adapter(path, adapter_name=name) |
|
|
| transformer.set_adapters(adapter_names, merge_weights) |
| transformer.requires_grad_(False) |
| transformer.eval() |
|
|
| model.transformer = transformer |
|
|
| if torch.cuda.is_available(): |
| mem_after = torch.cuda.memory_allocated(model.device) / 1024**2 |
| logger.info(f"[LoRA] GPU memory after load: {mem_after:.2f} MB") |
| |
| yield model |
|
|
| finally: |
| current_transformer = model.transformer |
| if isinstance(current_transformer, peft.PeftModel): |
| logger.info("[LoRA] Starting LoRA unload and cleanup...") |
| try: |
| unloaded_transformer = current_transformer.unload() |
| model.transformer = unloaded_transformer |
| except Exception as e: |
| logger.error(f"[LoRA] Failed to use unload(): {e}. Manually retrieving base model.") |
| model.transformer = current_transformer.base_model.model |
| |
| if hasattr(model.transformer, "peft_config"): |
| del model.transformer.peft_config |
| logger.info(f"[LoRA] Final cleanup: Removed residual peft_config from Base Model.") |
| |
| torch.cuda.empty_cache() |
|
|
| if torch.cuda.is_available(): |
| mem_final = torch.cuda.memory_allocated(model.device) / 1024**2 |
| logger.info(f"[LoRA] GPU memory after final unload: {mem_final:.2f} MB") |
|
|