File size: 8,754 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
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
import peft
import torch
import logging

from contextlib import contextmanager
from diffusers.pipelines.qwenimage.pipeline_qwenimage_edit_plus import QwenImageEditPlusPipeline, QwenImageEditPlusPipeline, QwenImageTransformer2DModel
from PIL import Image


COLOR_GRAY = "\033[37m"
COLOR_BLUE = "\033[94m"
COLOR_RED = "\033[31m"
COLOR_GREEN = "\033[32m"
COLOR_RESET = "\033[0m"


def get_logger(name: str, log_file: str):
    # 1. Get a logger instance
    logger = logging.Logger(name)

    # 2. Create a console handler
    console_handler = logging.StreamHandler()
    console_handler.setLevel(logging.INFO)  # Set the level for console output
    console_formatter = logging.Formatter(f"{COLOR_GREEN}[%(asctime)s][%(name)s][%(levelname)s]{COLOR_RESET} - %(message)s")
    console_handler.setFormatter(console_formatter)

    # 3. Create a file handler
    file_handler = logging.FileHandler(log_file)  # Specify the log file name
    file_handler.setLevel(logging.DEBUG)  # Set the level for file output (e.g., capture more details in the file)
    file_formatter = logging.Formatter("[%(asctime)s][%(name)s][%(levelname)s] - %(message)s")
    file_handler.setFormatter(file_formatter)

    # 4. Add handlers to the logger
    logger.addHandler(console_handler)
    logger.addHandler(file_handler)

    return logger


@contextmanager
def use_lora_adapter(
    model, # 应该是 QwenImageEditPlusPipeline 类型
    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]

    # --- 获取原始 Base Model 的引用 ---
    # 确保我们始终操作最底层的 Transformer 模块
    transformer = model.transformer
    while isinstance(transformer, peft.PeftModel):
        transformer = transformer.base_model.model
    
    # original_transformer = transformer # 存储一个干净的引用,以备后续清理和恢复
    
    # -----------------------------
    # 1. 彻底清理 Base Model (ENTER 阶段)
    # -----------------------------
    
    # 1a. 如果当前 Pipeline 中有 PeftModel,先卸载它
    if isinstance(model.transformer, peft.PeftModel):
        logger.info("[LoRA] Unloading existing PeftModel wrapper.")
        try:
            # unload() 会返回 Base Model,并清理 LoRA 权重
            model.transformer = model.transformer.unload()
        except AttributeError:
            # 兼容不同 peft 版本
            model.transformer = model.transformer.base_model.model
            
    # 1b. 【关键】从 Base Model 实例上强制清理残留的 PEFT 配置
    # Base Model 现在是 QwenImageTransformer2DModel 或其基类
    if hasattr(model.transformer, "peft_config"):
        # 尝试删除 PEFT 配置字典,这是配置残留的主要原因
        del model.transformer.peft_config
        logger.info(f"[LoRA] Successfully removed residual peft_config from Base Model.")
    
    # 确保 transformer 变量指向当前干净的 Base Model
    transformer = model.transformer 

    # -----------------------------
    # 2. Load new LoRA(s)
    # -----------------------------
    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")

        # 2a. 检查配置一致性并获取配置
        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]
        
        # 2b. 使用第一个适配器配置,将 Base Model 包装成 PeftModel
        # 注意:这里我们使用 adapter_names[0] 初始化第一个配置
        transformer = peft.get_peft_model(transformer, lora_config, adapter_name=adapter_names[0])
        logger.info(f"After get_peft_model: {type(transformer)=}")

        # 2c. 加载所有适配器(包括第一个,以确保权重加载)
        for path, name in zip(lora_paths, adapter_names):
            # 虽然第一个已初始化,但 load_adapter 确保权重被加载
            logger.info(f"[LoRA] Loading adapter: {name=}, {path=}")
            transformer.load_adapter(path, adapter_name=name)

        # 2d. 合并权重并设置模型状态
        # set_adapters 会设置 active_adapters 列表以及对应的加权系数
        transformer.set_adapters(adapter_names, merge_weights)
        transformer.requires_grad_(False)
        transformer.eval()

        # --- swap into pipeline ---
        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:
        # -----------------------------
        # 3. 彻底卸载 LoRA(s) (EXIT 阶段)
        # -----------------------------
        
        current_transformer = model.transformer

        if isinstance(current_transformer, peft.PeftModel):
            logger.info("[LoRA] Starting LoRA unload and cleanup...")
            # 3a. 卸载 PeftModel 包装器和权重
            try:
                # 这一步将清理 GPU 上的 LoRA 权重,并返回 Base Model
                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
        
        # 3b. 【关键】再次从 Base Model 上强制清理残留的 PEFT 配置
        # 确保 Base 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")


if __name__ == "__main__":

    device = "cuda"
    lora_paths = [
        "/temp/xr/StyQA-latency/lora_adapters/3D_Chibi/3D_Chibi_adapter",
    ]
    lora_paths2 = [
        "/temp/xr/StyQA-latency/lora_adapters/pixel_level_0/pixel_level_0_adapter",
        "/temp/xr/StyQA-latency/lora_adapters/pixel_level_1/pixel_level_1_adapter",
    ]
    adapter_names = [
        "3D_Chibi_adapter",
    ]
    adapter_names2 = [
        "pixel_level_0_adapter",
        "pixel_level_1_adapter",
    ]
    merge_weights = [1.0]
    merge_weights2 = [0.5, 0.5]
    model = QwenImageEditPlusPipeline.from_pretrained(
        "Qwen/Qwen-Image-Edit-2509",
        torch_dtype=torch.bfloat16,
    ).to(device)
    logger = get_logger(__name__, "lora_test.log")

    cnt_image = "/temp/rey/eval_datasets/unsplash_subset/_ZO2obWfpFQ.jpg"
    ref_image = "/temp/rey/eval_datasets/omniconsistency_subset/3D_Chibi/scene/040.jpg"
    ref_image2 = "/temp/rey/eval_datasets/style30k_subset/s0416____0912_01_query_2_img_000050_1683438561136_035474498831515366.jpg.jpg"
    limit = 3

    for i in range(limit):
        with use_lora_adapter(model, lora_paths, adapter_names, merge_weights, logger):
            output = model(
                image=Image.open(cnt_image).convert("RGB"),
                prompt="Modify the image to 3D Chibi style. Keep the content unchanged.",
                num_inference_steps=25,
                height=1024,
                width=1024,
                output_type="pil",
                return_dict=True,
                generator=torch.manual_seed(42),
            ).images[0]

            output.save(f"lora_test_output1_{i}.jpg")

        with use_lora_adapter(model, lora_paths2, adapter_names2, merge_weights2, logger):
            output = model(
                image=[Image.open(cnt_image).convert("RGB"), Image.open(ref_image2).convert("RGB")],
                prompt="",
                num_inference_steps=25,
                height=1024,
                width=1024,
                output_type="pil",
                return_dict=True,
                generator=torch.manual_seed(42),
            ).images[0]

            output.save(f"lora_test_output2_{i}.jpg")