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")