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, ): # normalize to list 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")