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