1.作者介绍

刘壮,男,西安工程大学电子信息学院,2025级研究生
研究方向:机器视觉与人工智能
电子邮箱:1523064593@qq.com

李逸超,男,西安工程大学电子信息学院,2025级研究生,张宏伟人工智能课题组

研究方向:机器视觉与人工智能

电子邮件:2317314922@qq.com

2.关于理论方面的知识介绍

2.1 图片提示词反推任务

图片提示词反推是指根据一张已经存在的图片,自动分析画面中的主体、场景、构图、光照、色彩、风格和细节,并生成一段可以再次用于文生图模型的提示词。它常用于 AI 绘画复刻、图像风格分析、素材整理和提示词学习。

本项目采用“视觉大模型理解 + sklearn 检索增强 + 规则化清洗输出”的组合方案。智谱视觉模型负责理解图片内容,sklearn 负责从参考提示词数据集中找出相似风格表达,最后由本地规则对标点、重复词和中英文输出格式进行整理。

图1  项目整体流程图

2.2 智谱AI视觉API调用原理

智谱 AI API 在项目中承担图像语义理解功能。程序首先读取本地图片,并将图片编码为 base64 字符串;随后把图片内容和提示语一起发送给视觉理解模型;模型返回对图片的自然语言描述。为了方便后续拼接提示词,项目要求模型尽量返回主体、场景、构图、光照、色彩、风格、情绪和可复用元素等字段。

API 调用的关键点包括:第一,API Key 从环境变量或 config.py 中读取,避免把敏感信息写死在主流程中;第二,图片需要转换为模型可接收的 data URL;第三,返回文本要经过结构化解析,才能稳定进入后续的提示词生成流程。

环节

说明

图片读取

读取本地 JPG/PNG 文件,检查文件是否存在

Base64 编码

将二进制图片转换为字符串,便于放入 API 请求体

视觉理解

调用智谱视觉模型,获得主体、场景、风格等描述

字段解析

把自然语言描述整理成字典字段,减少后续拼接混乱

异常处理

处理 API Key 缺失、网络异常、依赖库缺失等问题


2.3 sklearn最近邻检索算法

仅依靠视觉模型返回的描述,提示词有时会比较像普通说明文字,不够像绘图提示词。因此项目加入一个小型参考提示词数据集,并用 sklearn 的 StandardScaler 与 NearestNeighbors 进行检索增强。

程序会从输入图片中提取颜色均值、亮度、对比度、清晰度、边缘强度和尺寸比例等数值特征。参考数据集中每条样本也有对应的提示词与特征。经过标准化后,NearestNeighbors 根据欧氏距离寻找最相近的 Top-K 参考项,把其中的风格词加入最终反推结果。

图2  sklearn 最近邻检索增强原理

3.实验过程介绍

3.1 数据集与参考提示词库

本项目的数据集分为两部分:一部分是待反推的测试图片,例如 test1.jpg、test2.jpg;另一部分是自建参考提示词库,写在 reference_dataset.py 中。参考库规模较小,但覆盖了风景、花卉、人物等常见图像风格,适合作为课程项目展示 sklearn 检索流程。

 

图3  测试图片示例 test1.jpg test2.jpg

数据来源

数据内容

作用

本地测试图片

test1.jpg、test2.jpg 等

作为 API 视觉理解和提示词反推输入

reference_dataset.py

人工整理的参考提示词及风格标签

为 sklearn 最近邻检索提供可复用表达

image_features.py

从图片中提取颜色、亮度、清晰度等特征

把图片转换为可计算的数值向量

3.2 环境准备与运行方式

实验环境使用 ai_cpu 虚拟环境。运行前需要安装 zhipuai、scikit-learn、Pillow、numpy 等依赖,并配置智谱 API Key。项目支持在命令行中选择中文输出、英文输出或同时输出两种语言,方便多次测试同一张图片。

推荐运行命令如下:

cd /d D:\AI_Project\AI_homework\zhipu_image_prompt
python main.py --image test2.jpg

3.3代码模块设计


图4  项目代码模块结构图

main.py 是程序入口,负责解析图片路径并提供语言选择菜单;pipeline.py 是核心流程控制模块,负责串联图像理解、字段清洗、检索增强和最终提示词生成;zhipu_client.py 封装智谱 API;retriever.py 调用 sklearn 完成最近邻检索;image_features.py 负责将图片转换为机器学习可处理的特征向量。


3.4 完整代码展示
 

main.py:

from __future__ import annotations



import argparse

import json

import sys

from dataclasses import asdict

from pathlib import Path



# 兼容两种运行方式:

# 1. 在项目根目录执行:python zhipu_image_prompt/main.py

# 2. 在包目录中执行:python main.py

if __package__ in (None, ""):

    sys.path.insert(0, str(Path(__file__).resolve().parent.parent))



from zhipu_image_prompt.pipeline import ReversePromptPipeline, ReversePromptResult



def build_parser() -> argparse.ArgumentParser:

    """构建命令行参数解析器。"""



    parser = argparse.ArgumentParser(

        description="调用智谱 AI 和 sklearn 实现图片提示词反推。"

    )

    parser.add_argument("--image", required=True, help="输入图片路径。")

    parser.add_argument(

        "--top-k",

        type=int,

        default=3,

        help="sklearn 检索参考提示词的数量。",

    )

    parser.add_argument(

        "--json",

        action="store_true",

        help="以 JSON 格式一次性输出中英文完整结果。",

    )

    parser.add_argument("--name", default=None, help="汇报人姓名。")

    parser.add_argument("--student-id", default=None, help="汇报人学号。")

    return parser



def main() -> None:

    """命令行入口。"""



    parser = build_parser()

    args = parser.parse_args()



    pipeline = ReversePromptPipeline()

    report_name = args.name or pipeline.config.report_name

    student_id = args.student_id or pipeline.config.student_id



    # 缓存不同语言的结果,避免用户反复打印时重复调用智谱 API。

    cache: dict[str, ReversePromptResult] = {}



    def get_result(language: str) -> ReversePromptResult:

        if language not in cache:

            cache[language] = pipeline.run(

                image_path=args.image,

                top_k=args.top_k,

                language=language,

            )

        return cache[language]



    # JSON 模式适合调试或把结果交给其他程序处理。

    if args.json:

        payload = {

            "zh": asdict(get_result("zh")),

            "en": asdict(get_result("en")),

            "reporter": {

                "name": report_name,

                "student_id": student_id,

            },

        }

        print(json.dumps(payload, ensure_ascii=False, indent=2))

        return



    # 交互式菜单:运行一次程序,可以多次打印中文或英文结果。

    while True:

        print()

        print("请选择输出语言:")

        print("1. 中文反推提示词")

        print("2. English reverse prompt")

        print("3. 同时输出中文和英文")

        print("0. 退出")

        choice = input("请输入选项编号:").strip()



        if choice == "1":

            print_result(get_result("zh"), report_name, student_id)

        elif choice == "2":

            print_result(get_result("en"), report_name, student_id)

        elif choice == "3":

            print_result(get_result("zh"), report_name, student_id)

            print_result(get_result("en"), report_name, student_id)

        elif choice == "0":

            break

        else:

            print("无效选项,请输入 1、2、3 或 0。")



def print_result(

    result: ReversePromptResult,

    report_name: str,

    student_id: str,

) -> None:

    """打印最终反推结果。

    """



    if result.language == "en":

        print()

        print("=== Final Reverse Prompt ===")

        print(result.final_prompt)

        print()

        print("=== Reporter ===")

        print(f"Name: {report_name}")

        print(f"Student ID: {student_id}")

        return



    print()

    print("=== 最终反推结果 ===")

    print(result.final_prompt)

    print()

    print("=== 汇报信息 ===")

    print(f"姓名:{report_name}")

    print(f"学号:{student_id}")



if __name__ == "__main__":

    main()

pipeline.py:

from __future__ import annotations



from dataclasses import dataclass

from pathlib import Path

import re



from .config import load_config

from .image_features import load_image

from .retriever import PromptRetriever, RetrievalResult

from .zhipu_client import VisionUnderstanding, ZhipuImageAnalyzer



ZH_COMMA = ","

ZH_SPLIT_RE = r"[,,;;、。.!?!?]+"

EN_SPLIT_RE = r"[,;;,。]+"



@dataclass(slots=True)

class ReversePromptResult:

    """图片反推提示词流水线的最终结果。"""



    image_path: str

    language: str

    vision_understanding: VisionUnderstanding

    retrieval_results: list[RetrievalResult]

    reverse_keywords: list[str]

    cleaned_prompt_parts: list[str]

    final_prompt: str



class ReversePromptPipeline:

    """图片反推提示词总流程。



    该类负责串联四个阶段:

    1. 读取图片;

    2. 使用 sklearn 检索相似参考提示词;

    3. 调用智谱视觉模型理解图片;

    4. 清洗并融合关键词,得到最终反推提示词。

    """



    def __init__(self) -> None:

        self.config = load_config()

        self.retriever = PromptRetriever()

        self.analyzer = ZhipuImageAnalyzer(self.config)



    def run(

        self,

        image_path: str | Path,

        top_k: int = 3,

        language: str = "zh",

    ) -> ReversePromptResult:

        """执行完整反推流程。"""



        language = self._normalize_language(language)

        image_path = str(image_path)

        image = load_image(image_path)



        retrieval_results = self.retriever.retrieve(

            image=image,

            top_k=top_k,

            language=language,

        )

        vision_understanding = self.analyzer.analyze(

            image_path=image_path,

            language=language,

        )



        reverse_keywords = self._collect_parts(

            vision_understanding=vision_understanding,

            retrieved_prompt_parts=self._unique_retrieved_prompts(retrieval_results),

            language=language,

        )

        final_prompt = self._compose_final_prompt(reverse_keywords, language)



        return ReversePromptResult(

            image_path=image_path,

            language=language,

            vision_understanding=vision_understanding,

            retrieval_results=retrieval_results,

            reverse_keywords=reverse_keywords,

            cleaned_prompt_parts=reverse_keywords,

            final_prompt=final_prompt,

        )



    def _normalize_language(self, language: str) -> str:

        """统一语言参数,支持 zh/cn/chinese 和 en/english。"""



        language = language.lower().strip()

        if language in {"zh", "cn", "chinese"}:

            return "zh"

        if language in {"en", "english"}:

            return "en"

        raise ValueError("language must be 'zh' or 'en'")



    def _unique_retrieved_prompts(

        self,

        retrieval_results: list[RetrievalResult],

    ) -> list[str]:

        """去掉重复的检索参考提示词。"""



        prompts: list[str] = []

        seen_prompts: set[str] = set()

        for result in retrieval_results:

            if result.prompt not in seen_prompts:

                seen_prompts.add(result.prompt)

                prompts.append(result.prompt)

        return prompts



    def _collect_parts(

        self,

        vision_understanding: VisionUnderstanding,

        retrieved_prompt_parts: list[str],

        language: str,

    ) -> list[str]:

        """收集并清洗所有候选关键词。



        先加入智谱模型返回的结构化字段,再加入 reusable_elements,

        最后加入 sklearn 检索到的参考提示词片段。

        """



        parts: list[str] = []

        ordered_fields = [

            vision_understanding.subject,

            vision_understanding.scene,

            vision_understanding.description,

            vision_understanding.composition,

            vision_understanding.lighting,

            vision_understanding.color_palette,

            vision_understanding.style,

            vision_understanding.mood,

        ]



        for field in ordered_fields:

            self._append_phrase(parts, field, language)



        for item in vision_understanding.reusable_elements:

            self._append_phrase(parts, item, language)



        for retrieved_prompt in retrieved_prompt_parts:

            self._append_phrase(parts, retrieved_prompt, language)



        return parts



    def _append_phrase(self, parts: list[str], phrase: str, language: str) -> None:



        for chunk in self._split_raw_phrase(phrase, language):

            cleaned = self._clean_phrase(chunk, language)

            if not cleaned:

                continue



            normalized = cleaned.casefold()

            existing = {item.casefold() for item in parts}

            if normalized in existing:

                continue



            should_skip = False

            for item in parts:

                lower_item = item.casefold()

                if normalized in lower_item or lower_item in normalized:

                    should_skip = True

                    break



            if not should_skip:

                parts.append(cleaned)



    def _split_raw_phrase(self, phrase: object, language: str) -> list[str]:

        """根据语言拆分原始短语。"""



        text = self._normalize_raw_text(phrase, language)

        pattern = EN_SPLIT_RE if language == "en" else ZH_SPLIT_RE

        return [part for part in re.split(pattern, text) if part.strip()]



    def _normalize_raw_text(self, phrase: object, language: str) -> str:



        text = str(phrase or "")

        text = re.sub(r"[*#`_>]+", " ", text)

        text = re.sub(r"\s+", " ", text).strip()



        if language == "en":

            text = re.sub(r"[.。]+\s*,", ",", text)

            text = re.sub(r"\s*[,;、]\s*", ", ", text)

            text = re.sub(r"\s*,\s*", ", ", text)

            return text.strip(" ,;:。.")



        text = re.sub(r"\s*[;,,;、。.!?!?]\s*", ZH_COMMA, text)

        text = re.sub(rf"{ZH_COMMA}+", ZH_COMMA, text)

        return text.strip(f" {ZH_COMMA},;;、。.!?!?")



    def _clean_phrase(self, phrase: str, language: str) -> str:

        """清理单个短语两端残留的标点。"""



        text = self._normalize_raw_text(phrase, language)

        if language == "en":

            return text.strip(" ,;:.")



        text = re.sub(r"\s+", "", text)

        return text.strip(f"{ZH_COMMA},;;、。.!?!?")



    def _compose_final_prompt(self, keywords: list[str], language: str) -> str:

        """将清洗后的关键词组合为最终反推提示词。"""



        clean_keywords = [self._clean_phrase(keyword, language) for keyword in keywords]

        clean_keywords = [keyword for keyword in clean_keywords if keyword]

        delimiter = ", " if language == "en" else ZH_COMMA

        prompt = delimiter.join(clean_keywords)



        if language == "en":

            prompt = re.sub(r"[.。]+\s*,", ",", prompt)

            prompt = re.sub(r"(,\s*){2,}", ", ", prompt)

            return prompt.strip(" ,;:.")



        prompt = re.sub(rf"{ZH_COMMA}+", ZH_COMMA, prompt)

        prompt = re.sub(rf"{ZH_COMMA}\s*{ZH_COMMA}+", ZH_COMMA, prompt)

        return prompt.strip(f"{ZH_COMMA},;;、。.!?!?")


zhipu_client.py:

from __future__ import annotations



from dataclasses import dataclass

import json

import re



from .config import AppConfig

from .image_features import image_to_base64



ZH_COMMA = ","



@dataclass(slots=True)

class VisionUnderstanding:

    """智谱视觉模型返回的结构化图像理解结果。



    这些字段会被 pipeline 继续清洗、去重并组合成最终反推提示词。

    """



    description: str

    subject: str

    scene: str

    composition: str

    lighting: str

    color_palette: str

    style: str

    mood: str

    reusable_elements: list[str]

    raw_response: str



class ZhipuImageAnalyzer:

    """智谱视觉模型调用封装。



    该类只负责一件事:把本地图片和提示词指令发送给智谱 API,

    并将模型返回内容解析为 `VisionUnderstanding`。

    """



    def __init__(self, config: AppConfig) -> None:

        self.config = config



    def analyze(self, image_path: str, language: str = "zh") -> VisionUnderstanding:

        """分析图片并返回中文或英文结构化字段。"""



        if not self.config.api_key:

            raise ValueError(

                "ZHIPUAI_API_KEY is empty. Please set the environment variable first."

            )



        from zhipuai import ZhipuAI



        client = ZhipuAI(api_key=self.config.api_key)

        image_base64 = image_to_base64(image_path)



        # 智谱多模态接口的 content 中同时放入文本指令和图片 data URL。

        response = client.chat.completions.create(

            model=self.config.model,

            temperature=self.config.temperature,

            messages=[

                {

                    "role": "system",

                    "content": self._system_prompt(language),

                },

                {

                    "role": "user",

                    "content": [

                        {

                            "type": "text",

                            "text": self._user_prompt(language),

                        },

                        {

                            "type": "image_url",

                            "image_url": {

                                "url": f"data:image/png;base64,{image_base64}"

                            },

                        },

                    ],

                },

            ],

        )



        # SDK 可能返回字符串,也可能返回分段列表,因此先统一成字符串。

        content = self._normalize_message_content(response.choices[0].message.content)

        parsed = self._parse_json_payload(content)

        return VisionUnderstanding(

            description=self._clean_text(parsed.get("description", ""), language),

            subject=self._clean_text(parsed.get("subject", ""), language),

            scene=self._clean_text(parsed.get("scene", ""), language),

            composition=self._clean_text(parsed.get("composition", ""), language),

            lighting=self._clean_text(parsed.get("lighting", ""), language),

            color_palette=self._clean_text(parsed.get("color_palette", ""), language),

            style=self._clean_text(parsed.get("style", ""), language),

            mood=self._clean_text(parsed.get("mood", ""), language),

            reusable_elements=self._clean_list(parsed.get("reusable_elements", []), language),

            raw_response=content,

        )



    def _system_prompt(self, language: str) -> str:

        """构造 system prompt,用来限定模型角色和输出规则。"""



        if language == "en":

            return (

                "You are an image prompt reverse-engineering assistant. "

                "Analyze the image and extract prompt-ready information for image generation. "

                "Return valid JSON only. Do not use Markdown, code fences, or bullet lists."

            )

        return (

            "你是图片提示词反推助手。"

            "请根据输入图片提取适合文生图模型使用的中文提示词信息。"

            "只返回合法 JSON,不要使用 Markdown、代码块或项目符号。"

        )



    def _user_prompt(self, language: str) -> str:

        """构造 user prompt,指定返回字段和目标语言。"""



        keys = (

            "description, subject, scene, composition, lighting, "

            "color_palette, style, mood, reusable_elements"

        )

        if language == "en":

            return (

                f"Analyze this image in English and return JSON with these keys: {keys}. "

                "All values must be English. reusable_elements must be an array of short English phrases. "

                "Do not end field values with punctuation. Do not put sentence punctuation inside prompt phrases."

            )

        return (

            f"请用中文分析这张图片,并返回以下 JSON 字段:{keys}。"

            "所有字段值必须是中文。reusable_elements 必须是中文短语数组。"

            "不要在短语末尾添加句号、逗号或分号。"

        )



    def _normalize_message_content(self, message_content: object) -> str:

        """将 SDK 返回的 message.content 统一转换为纯文本。"""



        if isinstance(message_content, str):

            return message_content.strip()

        if isinstance(message_content, list):

            text_parts: list[str] = []

            for item in message_content:

                if isinstance(item, dict) and "text" in item:

                    text_parts.append(str(item["text"]))

                else:

                    text_parts.append(str(item))

            return "\n".join(part.strip() for part in text_parts if part.strip())

        return str(message_content).strip()



    def _parse_json_payload(self, content: str) -> dict[str, object]:

        """尽量从模型返回中解析出 JSON 对象。



        兜底策略:

        1. 直接解析完整返回;

        2. 如果被 Markdown 代码块包裹,先去掉代码块;

        3. 如果返回前后带了解释文字,就截取 `{...}` 部分;

        4. 如果仍然失败,就把整段内容作为 description 保存。

        """



        cleaned = content.strip()

        if cleaned.startswith("```"):

            cleaned = re.sub(r"^```[a-zA-Z0-9_-]*\s*", "", cleaned)

            cleaned = re.sub(r"\s*```$", "", cleaned)



        try:

            parsed = json.loads(cleaned)

            if isinstance(parsed, dict):

                return parsed

        except json.JSONDecodeError:

            pass



        match = re.search(r"\{.*\}", cleaned, flags=re.DOTALL)

        if match:

            try:

                parsed = json.loads(match.group(0))

                if isinstance(parsed, dict):

                    return parsed

            except json.JSONDecodeError:

                pass



        return self._fallback_parse(cleaned)



    def _fallback_parse(self, content: str) -> dict[str, object]:

        """模型没有按 JSON 返回时的最低可用结果。"""



        cleaned = self._clean_text(content, "zh")

        return {

            "description": cleaned,

            "subject": cleaned,

            "scene": "",

            "composition": "",

            "lighting": "",

            "color_palette": "",

            "style": "",

            "mood": "",

            "reusable_elements": [],

        }



    def _clean_text(self, value: object, language: str) -> str:

        """清理单个字段中的 Markdown 符号、重复空格和多余标点。"""



        delimiter = ", " if language == "en" else ZH_COMMA

        if isinstance(value, (list, tuple, set)):

            parts = [self._clean_text(item, language) for item in value]

            return delimiter.join(part for part in parts if part)



        text = str(value or "").strip()

        text = re.sub(r"[*#`_>]+", " ", text)

        text = re.sub(r"\s+", " ", text)



        if language == "en":

            text = re.sub(r"\s*[;;,、]\s*", ", ", text)

            text = re.sub(r"\s*,\s*", ", ", text)

            text = re.sub(r"[.。]+\s*,", ",", text)

            text = re.sub(r"[.。]+$", "", text)

            return text.strip(" ,;:.;,、。;")



        text = re.sub(r"\s*[;,,;、。.!?!?]\s*", ZH_COMMA, text)

        text = re.sub(rf"{ZH_COMMA}+", ZH_COMMA, text)

        return text.strip(f" {ZH_COMMA},;;、。.!?!?")



    def _clean_list(self, value: object, language: str) -> list[str]:

        """清洗数组字段,并按原顺序去重。"""



        if not isinstance(value, list):

            return []



        cleaned_items: list[str] = []

        seen: set[str] = set()

        for item in value:

            text = self._clean_text(item, language)

            if text and text not in seen:

                seen.add(text)

                cleaned_items.append(text)

        return cleaned_items


retriever.py:

from __future__ import annotations



from dataclasses import dataclass



from PIL import Image

from sklearn.neighbors import NearestNeighbors

from sklearn.preprocessing import StandardScaler



from .image_features import extract_feature_vector

from .reference_dataset import ReferenceSample, build_feature_matrix



@dataclass(slots=True)

class RetrievalResult:

    """一次 sklearn 近邻检索的命中结果。"""



    sample_name: str

    prompt: str

    distance: float



class PromptRetriever:

    """基于 sklearn NearestNeighbors 的提示词参考检索器。"""



    def __init__(self, n_neighbors: int = 3) -> None:

        # 构建参考库特征矩阵,每一行对应一个 ReferenceSample。

        feature_matrix, samples = build_feature_matrix()

        self.samples: list[ReferenceSample] = samples



        # 颜色直方图、亮度、宽高比等特征的量纲不同,先标准化再计算距离更合理。

        self.scaler = StandardScaler()

        scaled_matrix = self.scaler.fit_transform(feature_matrix)



        # 使用欧氏距离进行近邻检索。

        self.index = NearestNeighbors(

            n_neighbors=min(n_neighbors, len(samples)),

            metric="euclidean",

        )

        self.index.fit(scaled_matrix)



    def retrieve(

        self,

        image: Image.Image,

        top_k: int = 3,

        language: str = "zh",

    ) -> list[RetrievalResult]:

        """检索与输入图片最相似的 top_k 个参考样本。"""



        # 输入图像必须使用与参考库相同的特征提取和标准化流程。

        features = extract_feature_vector(image).reshape(1, -1)

        scaled_features = self.scaler.transform(features)

        distances, indices = self.index.kneighbors(

            scaled_features,

            n_neighbors=min(top_k, len(self.samples)),

        )



        results: list[RetrievalResult] = []

        for distance, idx in zip(distances[0], indices[0], strict=False):

            sample = self.samples[int(idx)]

            results.append(

                RetrievalResult(

                    sample_name=sample.name,

                    prompt=sample.prompt_for(language),

                    distance=float(distance),

                )

            )

        return results



image_features.py:
from __future__ import annotations



from pathlib import Path



import numpy as np

from PIL import Image, ImageFilter, ImageOps



def load_image(image_path: str | Path) -> Image.Image:

    """读取图片并统一转换为 RGB 格式。



    统一 RGB 后,后续特征提取不需要额外兼容灰度图、RGBA 图等不同格式。

    """



    return Image.open(image_path).convert("RGB")



def image_to_base64(image_path: str | Path) -> str:

    """将本地图片编码为 Base64 字符串。



    智谱视觉接口可以接收 data URL 形式的图片,因此这里先将文件内容转成

    Base64,再由 `zhipu_client.py` 拼接为 `data:image/png;base64,...`。

    """



    import base64



    with open(image_path, "rb") as file:

        return base64.b64encode(file.read()).decode("utf-8")



def extract_feature_vector(image: Image.Image) -> np.ndarray:

    """提取轻量级图像特征,用于 sklearn 相似图片检索。



    本项目没有使用复杂深度模型特征,而是使用传统统计特征,优点是:

    - 依赖少,容易运行;

    - 原理直观,适合课程汇报;

    - 可解释性强,能说明为什么两张图会被认为相似。



    特征组成:

    - RGB 三通道颜色直方图;

    - 平均亮度;

    - 对比度;

    - 边缘密度;

    - 饱和度代理值;

    - 宽高比。

    """



    # 缩放到固定尺寸,保证所有图片提取出的特征维度一致。

    resized = image.resize((224, 224))



    # 像素值归一化到 [0, 1],便于统计和距离计算。

    rgb = np.asarray(resized, dtype=np.float32) / 255.0



    # 分别统计 R/G/B 三个通道的颜色分布。

    # 每个通道 16 个 bins,总共得到 48 维颜色特征。

    hist_r, _ = np.histogram(rgb[:, :, 0], bins=16, range=(0.0, 1.0), density=True)

    hist_g, _ = np.histogram(rgb[:, :, 1], bins=16, range=(0.0, 1.0), density=True)

    hist_b, _ = np.histogram(rgb[:, :, 2], bins=16, range=(0.0, 1.0), density=True)



    # 灰度图用于计算整体亮度和对比度。

    gray = np.asarray(ImageOps.grayscale(resized), dtype=np.float32) / 255.0



    # FIND_EDGES 用于提取边缘,边缘越密集,说明轮廓或纹理越复杂。

    edges = np.asarray(

        ImageOps.grayscale(resized.filter(ImageFilter.FIND_EDGES)), dtype=np.float32

    ) / 255.0



    brightness = float(gray.mean())

    contrast = float(gray.std())

    edge_density = float((edges > 0.15).mean())

    saturation_proxy = float(np.std(rgb, axis=2).mean())

    aspect_ratio = float(image.width / max(image.height, 1))



    features = np.concatenate(

        [

            hist_r,

            hist_g,

            hist_b,

            np.array(

                [

                    brightness,

                    contrast,

                    edge_density,

                    saturation_proxy,

                    aspect_ratio,

                ],

                dtype=np.float32,

            ),

        ]

    )

    return features.astype(np.float32)


config.py:

from __future__ import annotations



import os

from dataclasses import dataclass



@dataclass(slots=True)

class AppConfig:

    """项目运行配置。



    这里集中保存智谱 API 调用和汇报输出需要的参数。

    使用 dataclass 可以让配置结构更清晰,也方便在 pipeline、client 等模块之间传递。

    """



    # 智谱开放平台 API Key。为空时不在配置阶段报错,而是在真正调用 API 时提示用户。

    api_key: str

    # 默认视觉模型名称,可通过环境变量 ZHIPU_MODEL 覆盖。

    model: str = "glm-4v-flash"

    # 模型采样温度。值越低越稳定,值越高越发散。

    temperature: float = 0.3



def load_config() -> AppConfig:

    """从环境变量读取配置。



    支持的环境变量:

    - ZHIPUAI_API_KEY:智谱 API Key

    - ZHIPU_MODEL:智谱视觉模型名

    - ZHIPU_TEMPERATURE:模型温度

    - REPORT_NAME:汇报人姓名

    - STUDENT_ID:汇报人学号

    """



    api_key = os.getenv("ZHIPUAI_API_KEY", "").strip()

    model = os.getenv("ZHIPU_MODEL", "glm-4v-flash").strip() or "glm-4v-flash"

    temperature_raw = os.getenv("ZHIPU_TEMPERATURE", "0.3").strip() or "0.3"

    report_name = os.getenv("REPORT_NAME", "刘壮").strip() or "刘壮"

    student_id = os.getenv("STUDENT_ID", "250412098").strip() or "250412098"



    try:

        temperature = float(temperature_raw)

    except ValueError as exc:

        raise ValueError(

            f"Invalid ZHIPU_TEMPERATURE value: {temperature_raw!r}"

        ) from exc



    return AppConfig(

        api_key=api_key,

        model=model,

        temperature=temperature,

        report_name=report_name,

        student_id=student_id,

    )


reference_dataset.py:
 

from __future__ import annotations



from dataclasses import dataclass



import numpy as np

from PIL import Image, ImageEnhance, ImageFilter

from sklearn.datasets import load_sample_images



from .image_features import extract_feature_vector



@dataclass(slots=True)

class ReferenceSample:

    """本地参考样本。



    每个样本包含:

    - name:样本名称,方便调试和展示;

    - prompt_zh:中文参考提示词;

    - prompt_en:英文参考提示词;

    - image:参考图片。

    """



    name: str

    prompt_zh: str

    prompt_en: str

    image: Image.Image



    def prompt_for(self, language: str) -> str:

        """根据语言选择返回中文或英文参考提示词。"""



        return self.prompt_en if language == "en" else self.prompt_zh



def _augment_image(image: Image.Image) -> dict[str, Image.Image]:

    """对参考图像做简单增强,扩充小型参考库。



    sklearn 内置示例图片数量很少,所以这里通过亮度、模糊、锐化和颜色增强

    生成几个变体,让近邻检索有更多可比较样本。

    """



    return {

        "original": image,

        "bright": ImageEnhance.Brightness(image).enhance(1.2),

        "soft": image.filter(ImageFilter.GaussianBlur(radius=1.2)),

        "sharp": ImageEnhance.Sharpness(image).enhance(1.8),

        "high_color": ImageEnhance.Color(image).enhance(1.3),

    }



def build_reference_samples() -> list[ReferenceSample]:

    """构建中英文参考提示词样本库。



    注意:这里的数据集不是用来训练模型,而是作为检索增强的参考库。

    输入图片会与这些参考图片计算相似度,然后借用相似样本绑定的提示词结构。

    """



    dataset = load_sample_images()

    images = dataset.images



    china = Image.fromarray(images[0]).convert("RGB")

    flower = Image.fromarray(images[1]).convert("RGB")



    sample_specs = {

        "china_landscape": {

            "image": china,

            "prompt_zh": (

                "纪实风景摄影,乡村道路,蓝天白云,绿色草地,远景透视,"

                "自然日光,真实色彩,空气感清晰,广角构图,高细节"

            ),

            "prompt_en": (

                "documentary landscape photography, rural road, blue sky, white clouds, "

                "green grass, distant perspective, natural daylight, realistic colors, "

                "clear atmosphere, wide-angle framing, high detail"

            ),

        },

        "flower_macro": {

            "image": flower,

            "prompt_zh": (

                "微距花卉摄影,近距离特写,花瓣纹理清晰,浅景深,柔和自然光,"

                "背景虚化,主体突出,色彩鲜明,高清细节"

            ),

            "prompt_en": (

                "macro flower photography, close-up subject, crisp petal texture, "

                "shallow depth of field, soft natural light, blurred background, "

                "strong subject separation, vivid color, high detail"

            ),

        },

    }



    samples: list[ReferenceSample] = []

    for base_name, spec in sample_specs.items():

        image = spec["image"]

        assert isinstance(image, Image.Image)

        for variant_name, variant_image in _augment_image(image).items():

            samples.append(

                ReferenceSample(

                    name=f"{base_name}_{variant_name}",

                    prompt_zh=str(spec["prompt_zh"]),

                    prompt_en=str(spec["prompt_en"]),

                    image=variant_image,

                )

            )

    return samples



def build_feature_matrix() -> tuple[np.ndarray, list[ReferenceSample]]:

    """将参考样本转换成 sklearn 可使用的二维特征矩阵。"""



    samples = build_reference_samples()

    matrix = np.vstack([extract_feature_vector(sample.image) for sample in samples])

    return matrix, samples



4.参考连接

1. 智谱 AI 开放平台文档:https://open.bigmodel.cn/

2. scikit-learn 官方文档:https://scikit-learn.org/stable/

3. Pillow 图像处理库文档:https://pillow.readthedocs.io/

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐