调用智谱AI实现实现图片反推提示词
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/
更多推荐




所有评论(0)