File size: 3,397 Bytes
59aed9d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
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")