StyQA / model /lora_utils.py
ReyChiaro's picture
Init commit
59aed9d
Raw
History Blame Contribute Delete
3.4 kB
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")