| 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): |
| |
| logger = logging.Logger(name) |
|
|
| |
| console_handler = logging.StreamHandler() |
| console_handler.setLevel(logging.INFO) |
| console_formatter = logging.Formatter(f"{COLOR_GREEN}[%(asctime)s][%(name)s][%(levelname)s]{COLOR_RESET} - %(message)s") |
| console_handler.setFormatter(console_formatter) |
|
|
| |
| file_handler = logging.FileHandler(log_file) |
| file_handler.setLevel(logging.DEBUG) |
| file_formatter = logging.Formatter("[%(asctime)s][%(name)s][%(levelname)s] - %(message)s") |
| file_handler.setFormatter(file_formatter) |
|
|
| |
| logger.addHandler(console_handler) |
| logger.addHandler(file_handler) |
|
|
| return logger |
|
|
|
|
| @contextmanager |
| def use_lora_adapter( |
| model, |
| lora_paths: str | list[str], |
| adapter_names: str | list[str], |
| merge_weights: float | list[float], |
| logger: logging.Logger, |
| ): |
| |
| 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") |
|
|
|
|
| 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") |
|
|