StyQA / lora_test.py
ReyChiaro's picture
Init commit
59aed9d
Raw
History Blame Contribute Delete
8.75 kB
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")