import os import json import torch from pathlib import Path from PIL import Image from typing import Literal from .data_struct import ( SYS_TASK_ROUTER, SYS_STYLE_ROUTER, SYS_ANALYSIS_SEMANTIC, SYS_ANALYSIS_PIXEL, SYS_TRANSFER, SYS_CRITERIA_CS, SYS_CRITERIA_RS, SYS_CRITERIA_DS, UserInput, SubTask, SubTaskOutput, AnalysisInput, AnalysisOutput, TransferInput, TransferOutput, CriteriaInput, CriteriaOutput, ) from .sub_tasks import AnalysisModule, TransferModule, CriteriaModule from model import QwenUMM, use_lora_adapter from utils import COLOR_GREEN, COLOR_RESET, PartialFormatter, load_image, get_logger, extract_json_result class StyQA: def __init__( self, output_dir: str = "agent_output", image_save_dir: str = "agent_output", log_dir: str = "agent_output", seed: int = 42, width: int = 1024, height: int = 1024, max_new_tokens: int = 1024, max_refine_times: int = 1, num_inference_steps: int = 16, lora_box: str = "prompts/lora_box.json", semantic_loras: str = "lora_adapters/semantic_loras.json", pixel_loras: str = "lora_adapters/pixel_loras.json", sys_prompt_dir: str = "prompts", device: str = "cuda:0", ): self.output_dir = output_dir self.image_save_dir = image_save_dir self.log_dir = log_dir self.seed = seed self.width = width self.height = height self.max_new_tokens = max_new_tokens self.max_refine_times = max_refine_times self.num_inference_steps = num_inference_steps self.device = device with open(lora_box) as f: self.lora_box = json.load(f) with open(semantic_loras) as f: self.semantic_loras = json.load(f) with open(pixel_loras) as f: self.pixel_loras = json.load(f) os.makedirs(self.output_dir, exist_ok=True) os.makedirs(self.image_save_dir, exist_ok=True) os.makedirs(self.log_dir, exist_ok=True) self.log_file = os.path.join(self.log_dir, "StyQA.log") self.logger = get_logger(__name__, self.log_file) title = "# ---- StyQA Configs ---- #" self.logger.info(title) self.logger.info(f"| {self.output_dir}") self.logger.info(f"| {self.image_save_dir}") self.logger.info(f"| {self.log_dir}") self.logger.info(f"| {self.seed}") self.logger.info(f"| {self.width}") self.logger.info(f"| {self.height}") self.logger.info(f"| {self.max_new_tokens}") self.logger.info(f"| {self.max_refine_times}") self.logger.info(f"| {self.num_inference_steps}") self.logger.info(f"| {self.device}") self.logger.info(f"| {self.log_file}") self.logger.info("# " + "-" * (len(title) - 4) + " #") sys_prompt_files = [f for f in os.listdir(sys_prompt_dir) if os.path.splitext(f)[1] == ".md"] for sys_prompt_file in sys_prompt_files: sys_prompt_name = os.path.basename(sys_prompt_file) with open(os.path.join(sys_prompt_dir, sys_prompt_file)) as f: setattr(self, f"SYS_{sys_prompt_name.upper()}", f.read()) self.logger.info(f"Load SYS_PROMPT: {sys_prompt_name}") if not hasattr(self, "SYS_TASK_ROUTER"): self.SYS_TASK_ROUTER = SYS_TASK_ROUTER if not hasattr(self, "SYS_STYLE_ROUTER"): lora_box_str = "" for k, v in self.lora_box.items(): desc = v["description"] lora_box_str += f"- {k}: {desc}\n" self.SYS_STYLE_ROUTER = SYS_STYLE_ROUTER.format_map(PartialFormatter(style_value_and_descriptions=lora_box_str)) if not hasattr(self, "SYS_ANALYSIS_SEMANTIC"): self.SYS_ANALYSIS_SEMANTIC = SYS_ANALYSIS_SEMANTIC if not hasattr(self, "SYS_ANALYSIS_PIXEL"): self.SYS_ANALYSIS_PIXEL = SYS_ANALYSIS_PIXEL if not hasattr(self, "SYS_TRANSFER"): self.SYS_TRANSFER = SYS_TRANSFER if not hasattr(self, "SYS_CRITERIA_CS"): self.SYS_CRITERIA_CS = SYS_CRITERIA_CS if not hasattr(self, "SYS_CRITERIA_RS"): self.SYS_CRITERIA_RS = SYS_CRITERIA_RS if not hasattr(self, "SYS_CRITERIA_DS"): self.SYS_CRITERIA_DS = SYS_CRITERIA_DS # -- Load Model -- # model_title = "# ---- Load Model ---- #" self.logger.info(model_title) self.model = QwenUMM(device=self.device) self.logger.info(f"# " + "-" * (len(model_title) - 4) + " #") # -- Init Task Modules -- # self.analysis_module = AnalysisModule(self.max_new_tokens) self.transfer_module = TransferModule( num_inference_steps=self.num_inference_steps, height=self.height, width=self.width, seed=self.seed, ) self.criteria_module = CriteriaModule(self.max_new_tokens) # -------------------------------- # # -------- Helper Methods -------- # # -------------------------------- # def task_router(self, user_input: UserInput) -> list[SubTask]: title = "# ---- Task Router ---- #" self.logger.info(title) prompt = user_input.prompt ref_dict = user_input.ref_dict # -- Prompt split and Workflow extract -- # # Split prompt into picture centric # Each sub-prompt corresponds to one reference image # Extract the style transfer workflow from prompt # Output: List of JSON, JSON keys: ref_id, ref_prompt output = self.model( task="txt-gen", image=None, prompt=prompt, sys_prompt=self.SYS_TASK_ROUTER, max_new_tokens=self.max_new_tokens, ) self.logger.debug(f"Model raw output:\n{output}\n") output = extract_json_result(output, self.logger) if isinstance(output, dict): output = [output] if isinstance(output, list): _test = output[0] if "raw_output" in _test.keys(): output = [{"ref_id": 1, "ref_prompt": _test["raw_output"]}] self.logger.info(f"{COLOR_GREEN}Extract output:\n{output}\n{COLOR_RESET}") # -- Define subtasks -- # # Style Task: For each reference image, detect use semantic or pixel. # Style Value: For each reference image, extract the features or stylization strength. task_pipeline = [] for i, ref_item in enumerate(output): ref_key = f"Picture {ref_item['ref_id']}" ref_image_path = ref_dict[ref_key] style_info = self.model( task="txt-gen", image={"Picture 1": load_image(ref_image_path)}, prompt=ref_item["ref_prompt"], sys_prompt=self.SYS_STYLE_ROUTER, max_new_tokens=self.max_new_tokens, ) self.logger.debug(f"Model raw output for reference Picture {i+1}: \n{style_info}") style_info = extract_json_result(style_info, self.logger) self.logger.info(f"{COLOR_GREEN}Extract output:\n{style_info}{COLOR_RESET}") sub_task = SubTask( ref_id=ref_item["ref_id"], ref_image_path=ref_dict[ref_key], style_type=style_info["style_type"], style_value=style_info["style_value"], ) task_pipeline.append(sub_task) self.logger.info("# " + "-" * (len(title) - 4) + " #") return task_pipeline def optimize_instruction(self, prompt: str, cnt_image_or_path: str | Image.Image) -> str: title = "# ---- Optimize Instruction ---- #" self.logger.info(title) # 1. Detect objects in cnt image detect_output = self.model( task="txt-gen", image={"Picture 1": load_image(cnt_image_or_path)}, prompt="Detect the contents/objects/subjects in a list format, without explanations.", sys_prompt="", max_new_tokens=self.max_new_tokens, ) # 2. Generate instructions prompt = f"Style Description: {prompt}\nObject List:{detect_output}" instructions = self.model( task="txt-gen", image=None, prompt=prompt, sys_prompt=r"""You are a style-transfer expert. Your task is to apply a given style description to all objects in a provided list, ensuring that each object adopts the same style characteristics. Output the results as a list of instruction-style modifications, describing how each object should be transformed to match the target style.""", max_new_tokens=self.max_new_tokens, ) self.logger.info("# " + "-" * (len(title) - 4) + " #") return instructions # ----------------------------------- # # -------- Input Constructor -------- # # ----------------------------------- # def create_analysis_input(self, sub_task: SubTask, suggestion: str = "") -> AnalysisInput: sys_prompt = self.SYS_ANALYSIS_SEMANTIC if sub_task.style_type == "semantic" else self.SYS_ANALYSIS_PIXEL # For analysis, the reference images are handled # sequentially, the ref_id is not matter analysis_input = AnalysisInput( ref_id=sub_task.ref_id, ref_image_or_path=sub_task.ref_image_path, style_type=sub_task.style_type, style_value=sub_task.style_value, suggestion=suggestion, sys_prompt=sys_prompt, ) return analysis_input def create_transfer_input( self, instruct: str, cnt_image_or_path: str | Image.Image, sub_task: SubTask, suggestion: str = "", ) -> TransferInput: sys_prompt = self.SYS_TRANSFER transfer_input = TransferInput( prompt=instruct + f"\nSuggestion: {suggestion}", cnt_image_or_path=cnt_image_or_path, ref_image_or_path=sub_task.ref_image_path if sub_task.style_type != "semantic" else None, sys_prompt=sys_prompt, ) return transfer_input def create_criteria_input( self, instruction: str, cnt_image_or_path: str | Image.Image, sty_image_or_path: str | Image.Image, sub_task: SubTask, ) -> CriteriaInput: sys_prompts = { "cs": self.SYS_CRITERIA_CS, "rs": self.SYS_CRITERIA_RS, "ds": self.SYS_CRITERIA_DS, } criteria_input = CriteriaInput( instruction=instruction, cnt_image_or_path=cnt_image_or_path, ref_image_or_path=sub_task.ref_image_path, sty_image_or_path=sty_image_or_path, sys_prompts=sys_prompts, ) return criteria_input # ---------------------------- # # -------- LoRA tools -------- # # ---------------------------- # def config_lora_adapter(self, style_type: Literal["semantic", "pixel"], style_value: str | float) -> callable: title = "# ---- Config LoRA Adapters ---- #" self.logger.info(title) if style_type == "semantic": if style_value in self.semantic_loras.keys(): # Load pre-defined style type LoRA adapter lora_paths = [self.semantic_loras[style_value]["path"]] adapter_names = [self.semantic_loras[style_value]["adapter_name"]] merge_weight = [1.0] elif style_type == "pixel": # For pixel level <= 0: load level_0 adapter pixel_level = [0] # For pixel level in (0, 0.5): load level_0 and level_1 adapters and merge pixel_level = [0, 1] if 0 < style_value < 0.5 else pixel_level # For pixel level == 0.5: load level_1 adapter pixel_level = [1] if style_value == 0.5 else pixel_level # For pixel level in (0.5, 1.0): load level_2 adapter pixel_level = [1, 2] if 0.5 < style_value < 1.0 else pixel_level # For pixel level >= 1.0: load level_2 adapter pixel_level = [2] if style_value >= 1.0 else pixel_level lora_key = [f"level_{i}" for i in pixel_level] lora_paths = [self.pixel_loras[k]["path"] for k in lora_key] adapter_names = [self.pixel_loras[k]["adapter_name"] for k in lora_key] merge_weight = [1.0] if len(pixel_level) == 2: if 0 < style_value < 0.5: merge_weight = [(1.0 - style_value * 2), style_value * 2] elif 0.5 < style_value < 1.0: merge_weight = [(1.0 - (style_value - 0.5) * 2), (style_value - 0.5) * 2] else: self.logger.info(f"No suitable LoRA adapter find for {style_type=}, {style_value=}") self.logger.info(f"{adapter_names=}, {lora_paths=}, {merge_weight=}") self.logger.info("# " + "-" * (len(title) - 4) + " #") return lora_paths, adapter_names, merge_weight # ------------------------------------ # # -------- Single Task Runner -------- # # ------------------------------------ # def run_analysis(self, analysis_input: AnalysisInput) -> AnalysisOutput: return self.analysis_module.run(self.model, analysis_input, self.logger) def run_transfer(self, transfer_input: TransferInput) -> TransferOutput: return self.transfer_module.run(self.model, transfer_input, self.logger) def run_criteria(self, criteria_input: CriteriaInput) -> CriteriaOutput: return self.criteria_module.run(self.model, criteria_input, self.logger) # -------------------------------------- # # -------- Composed Task Runner -------- # # -------------------------------------- # def run_analysis_to_optim_instruct( self, style_type: Literal["semantic", "pixel"], style_values: list[str | float], cnt_image_paths: list[str], ref_image_paths: list[str], ): """ Used to generate instructions based on analysis results. Return liset of instructions for `style_type` and `style_value`. """ items_to_save = [] for cnt_image_path, ref_image_path, style_value in zip(cnt_image_paths, ref_image_paths, style_values): instruction = "" item_to_save = {} item_to_save["content"] = cnt_image_path item_to_save["style"] = ref_image_path sub_task = SubTask( ref_id=1, ref_image_path=ref_image_path, style_type=style_type, style_value=style_value, ) analysis_input = self.create_analysis_input(sub_task, "") analysis_output = self.run_analysis(analysis_input) style_desc = getattr(analysis_output, style_type) instruction = "Transfer the style of Picture 1 into target style. The style is:\n" for k, v in style_desc.items(): instruction += f"{k}: {v}" item_to_save["description"] = instruction instruction = self.optimize_instruction(instruction, cnt_image_path) item_to_save["instruction"] = instruction if isinstance(style_value, str): item_to_save["category"] = style_value items_to_save.append(item_to_save) return items_to_save def run_transfer_with_lora( self, style_type: Literal["semantic", "pixel"], style_value: str | float, cnt_image_paths: list[str], ref_image_paths: list[str], enable_analysis: bool = True, convert_instruct: bool = False, save_dir: str = "", image_name_fmt="{cnt_image_name}@{ref_image_name}.jpg", ): lora_paths, adapter_names, merge_weight = self.config_lora_adapter(style_type, style_value) # unload_lora_adapters_fn = self.load_lora_adapter(style_type, style_value) image_save_dir = os.path.join(save_dir, "images") record_file = os.path.join(save_dir, "log_StyQA.jsonl") os.makedirs(image_save_dir, exist_ok=True) self.logger.info(self.model.model.transformer.active_adapters) with use_lora_adapter(self.model.model, lora_paths, adapter_names, merge_weight, self.logger): # -- Analysis -- # for cnt_image_path, ref_image_path in zip(cnt_image_paths, ref_image_paths): instruction = "" analysis_elapsed_sec = 0 item_to_save = {} item_to_save["content"] = cnt_image_path item_to_save["style"] = ref_image_path sub_task = SubTask( ref_id=1, ref_image_path=ref_image_path, style_type=style_type, style_value=style_value, ) if enable_analysis: analysis_start_event = torch.cuda.Event(enable_timing=True) analysis_end_event = torch.cuda.Event(enable_timing=True) torch.cuda.synchronize() analysis_start_event.record() analysis_input = self.create_analysis_input(sub_task, "") analysis_output = self.run_analysis(analysis_input) style_desc = getattr(analysis_output, style_type) instruction = "Transfer the style of Picture 1 into target style. The style is:\n" for k, v in style_desc.items(): instruction += f"{k}: {v}" if convert_instruct: instruction = self.optimize_instruction(instruction, cnt_image_path) analysis_end_event.record() torch.cuda.synchronize() analysis_elapsed_sec = analysis_start_event.elapsed_time(analysis_end_event) / 1000 item_to_save["instruction"] = instruction item_to_save["analysis_elapsed_sec"] = analysis_elapsed_sec transfer_start_event = torch.cuda.Event(enable_timing=True) transfer_end_event = torch.cuda.Event(enable_timing=True) torch.cuda.synchronize() transfer_start_event.record() transfer_input = self.create_transfer_input( instruct=( f"Transfer the style of Picture 1 to the style of Picture 2.\n" if not instruction else instruction ), # instruction, cnt_image_or_path=cnt_image_path, sub_task=sub_task, suggestion="", ) transfer_output = self.run_transfer(transfer_input) transfer_end_event.record() torch.cuda.synchronize() transfer_elapsed_sec = transfer_start_event.elapsed_time(transfer_end_event) / 1000 item_to_save["transfer_elapsed_sec"] = transfer_elapsed_sec item_to_save["elapsed_sec"] = analysis_elapsed_sec + transfer_elapsed_sec sty_image = transfer_output.sty_image save_name = image_name_fmt.format( cnt_image_name=Path(cnt_image_path).stem, ref_image_name=Path(ref_image_path).stem, ) if isinstance(style_value, str): save_name = Path(save_name).stem + f"@{style_value}.jpg" save_path = os.path.join(image_save_dir, save_name) item_to_save["output"] = save_path self.logger.info(f"Stylized image saved to {save_path}") sty_image.save(save_path) with open(record_file, "a") as f: f.write(json.dumps(item_to_save) + "\n") def run_transfer_without_lora( self, style_type: Literal["semantic", "pixel"], style_value: str | float, cnt_image_paths: list[str], ref_image_paths: list[str], enable_analysis: bool = True, convert_instruct: bool = False, save_dir: str = "", image_name_fmt="{cnt_image_name}@{ref_image_name}.jpg", ): # lora_paths, adapter_names, merge_weight = self.config_lora_adapter(style_type, style_value) # unload_lora_adapters_fn = self.load_lora_adapter(style_type, style_value) image_save_dir = os.path.join(save_dir, "images") record_file = os.path.join(save_dir, "log_StyQA.jsonl") os.makedirs(image_save_dir, exist_ok=True) self.logger.info(self.model.model.transformer.active_adapters) # with use_lora_adapter(self.model.model, lora_paths, adapter_names, merge_weight, self.logger): # -- Analysis -- # for cnt_image_path, ref_image_path in zip(cnt_image_paths, ref_image_paths): instruction = "" analysis_elapsed_sec = 0 item_to_save = {} item_to_save["content"] = cnt_image_path item_to_save["style"] = ref_image_path sub_task = SubTask( ref_id=1, ref_image_path=ref_image_path, style_type=style_type, style_value=style_value, ) if enable_analysis: analysis_start_event = torch.cuda.Event(enable_timing=True) analysis_end_event = torch.cuda.Event(enable_timing=True) torch.cuda.synchronize() analysis_start_event.record() analysis_input = self.create_analysis_input(sub_task, "") analysis_output = self.run_analysis(analysis_input) style_desc = getattr(analysis_output, style_type) instruction = "Transfer the style of Picture 1 into target style. The style is:\n" for k, v in style_desc.items(): instruction += f"{k}: {v}" if convert_instruct: instruction = self.optimize_instruction(instruction, cnt_image_path) analysis_end_event.record() torch.cuda.synchronize() analysis_elapsed_sec = analysis_start_event.elapsed_time(analysis_end_event) / 1000 item_to_save["instruction"] = instruction item_to_save["analysis_elapsed_sec"] = analysis_elapsed_sec transfer_start_event = torch.cuda.Event(enable_timing=True) transfer_end_event = torch.cuda.Event(enable_timing=True) torch.cuda.synchronize() transfer_start_event.record() transfer_input = self.create_transfer_input( instruct=( f"Transfer the style of Picture 1 to the style of Picture 2.\n" if not instruction else instruction ), # instruction, cnt_image_or_path=cnt_image_path, sub_task=sub_task, suggestion="", ) transfer_output = self.run_transfer(transfer_input) transfer_end_event.record() torch.cuda.synchronize() transfer_elapsed_sec = transfer_start_event.elapsed_time(transfer_end_event) / 1000 item_to_save["transfer_elapsed_sec"] = transfer_elapsed_sec item_to_save["elapsed_sec"] = analysis_elapsed_sec + transfer_elapsed_sec sty_image = transfer_output.sty_image save_name = image_name_fmt.format( cnt_image_name=Path(cnt_image_path).stem, ref_image_name=Path(ref_image_path).stem, ) if isinstance(style_value, str): save_name = Path(save_name).stem + f"@{style_value}.jpg" save_path = os.path.join(image_save_dir, save_name) item_to_save["output"] = save_path self.logger.info(f"Stylized image saved to {save_path}") sty_image.save(save_path) with open(record_file, "a") as f: f.write(json.dumps(item_to_save) + "\n") # --------------------------------- # # -------- Pipeline Runner -------- # # --------------------------------- # def run_subtask( self, cnt_image_path: str, iter_cnt_image: Image.Image, sub_task: SubTask, suggestion: str, convert_instruct: bool = True, refine_iter: int = 0, image_name_fmt="{cnt_image_name}@{ref_image_name}@iter{refine_iter}.jpg", ) -> SubTaskOutput: style_type = sub_task.style_type style_value = sub_task.style_value # -- 1. Style Analysis -- # analysis_input = self.create_analysis_input(sub_task, suggestion) analysis_output = self.run_analysis(analysis_input) style_desc: dict[str, str] = getattr(analysis_output, style_type) # -- [Optional] Optimize to Instructions -- # style_instruct = "" if convert_instruct: style_desc_str = "" for k, v in style_desc.items(): style_desc_str += f"- {k}: {v}\n" style_instruct = self.optimize_instruction(style_desc_str, cnt_image_path) # -- 2. Style Transfer -- # lora_paths, adapter_names, merge_weights = self.config_lora_adapter(style_type, style_value) with use_lora_adapter(self.model.model, lora_paths, adapter_names, merge_weights, self.logger): transfer_input = self.create_transfer_input(style_instruct, iter_cnt_image, sub_task, suggestion) transfer_output = self.run_transfer(transfer_input) sty_image = transfer_output.sty_image # Save stylized image save_image_name = image_name_fmt.format( cnt_image_name=Path(cnt_image_path).stem, ref_image_name=Path(sub_task.ref_image_path).stem, refine_iter=refine_iter, ) save_image_path = os.path.join(self.image_save_dir, save_image_name) self.logger.info(f"Stylized image saved to {save_image_path}") sty_image.save(save_image_path) self.logger.info(f"Transfer finished, image saved to: {save_image_path}") # -- 3. Style Criteria -- # # Always use cnt_image_path, not the iter_cnt_image criteria_input = self.create_criteria_input(style_instruct, cnt_image_path, sty_image, sub_task) criteria_output = self.run_criteria(criteria_input) return SubTaskOutput( cnt_image_path=cnt_image_path, ref_image_path=sub_task.ref_image_path, analysis_output=analysis_output, transfer_output=transfer_output, criteria_output=criteria_output, ) def run_pipeline(self, user_input: UserInput, update_suggestion: bool = False): title = "# ---- Run Pipeline ---- #" self.logger.info(title) sub_tasks = self.task_router(user_input) suggestions = [""] * len(sub_tasks) for refine_count in range(self.max_refine_times): refine_title = f"| ---- Refine [{refine_count+1}/{self.max_refine_times}] ---- |" self.logger.info(refine_title) # -- Prepare args used for iterations -- # iter_cnt_image = load_image(user_input.cnt_image_path) for i, sub_task in enumerate(sub_tasks): # Run for a single subtask output: SubTaskOutput = self.run_subtask( cnt_image_path=user_input.cnt_image_path, iter_cnt_image=iter_cnt_image, sub_task=sub_task, suggestion=suggestions[i], convert_instruct=True, refine_iter=i, image_name_fmt="{cnt_image_name}@{ref_image_name}@iter{refine_iter}.jpg", ) # After every subtask finished, the cnt image should be updated iter_cnt_image = output.transfer_output.sty_image # The suggestion or the stylization strength should be updated # Update suggestions for semantic task as it is instruction-motivated task if update_suggestion: if sub_task.style_type == "semantic": suggestions[i] = f"Content preservation: {output.criteria_output.cs_score['suggestion']}\n" suggestions[i] += f"Instruction following: {output.criteria_output.ds_score['suggestion']}" # Update stylization strength if sub_task.style_type == "pixel": if output.criteria_output.rs_score["suggestion"] == "increase": sub_tasks[i].style_value = min(1.0, sub_tasks[i].style_value + 0.1) elif output.criteria_output.rs_score["suggestion"] == "decrease": sub_tasks[i].style_value = max(0.0, sub_tasks[i].style_value - 0.1) # else: Unchanged self.logger.info("| " + "-" * (len(refine_title) - 4) + " |") self.logger.info("# " + "-" * (len(title) - 4) + " #")