kit初版,模型引入,agent优化
This commit is contained in:
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,15 +1,19 @@
|
||||
"""Agent Loop 主循环:LLM 对话 -> 工具调用 -> 结果回传 -> 继续。"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import AsyncGenerator, Optional
|
||||
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from app.agent.tools import TOOL_DEFINITIONS, execute_tool
|
||||
from app.config import get_llm_model, get_llm_max_iterations
|
||||
from app.services.image_gen import to_data_uri
|
||||
|
||||
|
||||
def _get_client() -> AsyncOpenAI:
|
||||
return AsyncOpenAI()
|
||||
base_url = os.getenv("OPENAI_BASE_URL")
|
||||
return AsyncOpenAI(base_url=base_url) if base_url else AsyncOpenAI()
|
||||
|
||||
SYSTEM_PROMPT = """\
|
||||
你是一个专业的游戏美术 AI 助手。你的工作是帮助美术人员通过对话生成游戏美术资源。
|
||||
@@ -19,6 +23,7 @@ SYSTEM_PROMPT = """\
|
||||
- 理解用户的审美意图,将中文描述转化为高质量的英文生成 prompt
|
||||
- 根据用户反馈迭代修改(调整颜色、风格、构图等)
|
||||
- 如果用户提供了参考图,将参考图的风格元素融入生成 prompt
|
||||
- 理解用户在图片上的标注(框选区域 + 文字批注),精准定位需要修改的部分
|
||||
|
||||
## 工作流程
|
||||
1. 理解用户需求,必要时追问细节(尺寸、风格、用途等)
|
||||
@@ -26,11 +31,20 @@ SYSTEM_PROMPT = """\
|
||||
3. 向用户展示结果并询问反馈
|
||||
4. 根据反馈调整 prompt 并重新生成
|
||||
|
||||
## 标注理解
|
||||
当用户发送带有「图片标注」的消息时,表示用户在之前生成的图片上做了标注。
|
||||
标注格式为:`[区域 (x%, y%) 大小 w%×h%]: 修改意见`
|
||||
- 区域坐标表示标注框在图片上的相对位置
|
||||
- 你需要理解标注区域所指的图片内容,并将修改意见融入新的 prompt
|
||||
- 如果附带了标注截图(参考图),仔细观察红色标注框和文字来理解用户意图
|
||||
- 在调整 prompt 时,保持原图整体风格不变,只针对标注区域做修改
|
||||
|
||||
## 生成 prompt 要求
|
||||
- 必须使用英文
|
||||
- 尽量详细描述:主体内容、风格、颜色方案、光照、构图、材质等
|
||||
- 尽量详细描述:主体内容、颜色方案、光照、构图、材质等
|
||||
- 如果用户要求游戏 UI 元素,添加相关关键词如 "game UI", "icon", "button" 等
|
||||
- 如果用户提供了参考图,在 prompt 中描述参考图的风格特征
|
||||
- 如果你能直接看到参考图(图片内容),可以在 prompt 中描述参考图的风格特征
|
||||
- 如果你无法看到参考图(只收到了文字提示说有参考图),参考图会由生图工具的 IP-Adapter 自动处理风格融合。此时你的 prompt 中**绝对不要猜测或指定任何画风/艺术风格关键词**(如 pixel art、watercolor、oil painting 等),只描述画面内容(主体、构图、光照等),把风格完全交给参考图来决定
|
||||
|
||||
## 注意事项
|
||||
- 用中文和用户交流
|
||||
@@ -42,6 +56,7 @@ SYSTEM_PROMPT = """\
|
||||
async def run_agent_loop(
|
||||
messages: list[dict],
|
||||
ref_image_url: Optional[str] = None,
|
||||
image_model: Optional[str] = None,
|
||||
) -> AsyncGenerator[dict, None]:
|
||||
"""
|
||||
运行 Agent Loop,以 SSE 事件流形式 yield 结果。
|
||||
@@ -56,33 +71,45 @@ async def run_agent_loop(
|
||||
# 构建消息列表
|
||||
api_messages = [{"role": "system", "content": SYSTEM_PROMPT}]
|
||||
|
||||
# 检测当前 LLM 是否支持 vision(多模态图片输入)
|
||||
llm_model = get_llm_model().lower()
|
||||
vision_capable = any(kw in llm_model for kw in ("gpt-4o", "gpt-4-vision", "claude"))
|
||||
|
||||
for msg in messages:
|
||||
if msg["role"] == "user" and ref_image_url:
|
||||
# 最后一条用户消息附加参考图(GPT-4o vision)
|
||||
if msg == messages[-1] or (
|
||||
msg.get("role") == "user"
|
||||
and messages.index(msg) == len(messages) - 1
|
||||
):
|
||||
if msg["role"] == "user" and ref_image_url and msg is messages[-1]:
|
||||
if vision_capable:
|
||||
# 支持 vision 的模型:直接发送图片
|
||||
api_messages.append({
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": msg["content"]},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": ref_image_url},
|
||||
"image_url": {"url": to_data_uri(ref_image_url)},
|
||||
},
|
||||
],
|
||||
})
|
||||
continue
|
||||
else:
|
||||
# 不支持 vision 的模型:以文字提示告知有参考图,风格由 IP-Adapter 处理
|
||||
hint = (
|
||||
f"{msg['content']}\n\n"
|
||||
"【系统提示:用户上传了一张参考图,已自动传递给图片生成工具的 IP-Adapter。"
|
||||
"IP-Adapter 会从参考图中提取风格并融合到生成结果中。"
|
||||
"你无法看到这张参考图,因此在生成 prompt 时:\n"
|
||||
"1. 只描述画面内容(主体、构图、光照、材质等)\n"
|
||||
"2. 不要猜测或添加任何画风/艺术风格关键词(如 pixel art、watercolor、cartoon 等)\n"
|
||||
"3. 风格完全由参考图通过 IP-Adapter 决定】"
|
||||
)
|
||||
api_messages.append({"role": "user", "content": hint})
|
||||
continue
|
||||
api_messages.append({"role": msg["role"], "content": msg["content"]})
|
||||
|
||||
client = _get_client()
|
||||
|
||||
max_iterations = 5
|
||||
for _ in range(max_iterations):
|
||||
for _ in range(get_llm_max_iterations()):
|
||||
try:
|
||||
response = await client.chat.completions.create(
|
||||
model="gpt-4o-mini",
|
||||
model=get_llm_model(),
|
||||
messages=api_messages,
|
||||
tools=TOOL_DEFINITIONS,
|
||||
stream=True,
|
||||
@@ -163,13 +190,28 @@ async def run_agent_loop(
|
||||
except json.JSONDecodeError:
|
||||
arguments = {}
|
||||
|
||||
result = await execute_tool(tool_name, arguments, ref_image_url)
|
||||
result = await execute_tool(tool_name, arguments, ref_image_url, image_model)
|
||||
|
||||
used_model = result.get("model_name", "")
|
||||
|
||||
if result.get("errors"):
|
||||
yield {
|
||||
"type": "tool_error",
|
||||
"data": {
|
||||
"tool": tool_name,
|
||||
"errors": result["errors"],
|
||||
"model_name": used_model,
|
||||
},
|
||||
}
|
||||
|
||||
# 如果有图片结果,推送给前端
|
||||
if result.get("images"):
|
||||
yield {
|
||||
"type": "image_result",
|
||||
"data": {"images": result["images"]},
|
||||
"data": {
|
||||
"images": result["images"],
|
||||
"prompt_used": result.get("prompt_used", ""),
|
||||
"model_name": used_model,
|
||||
},
|
||||
}
|
||||
|
||||
# 工具结果回传给 LLM
|
||||
@@ -179,7 +221,7 @@ async def run_agent_loop(
|
||||
"content": json.dumps(result, ensure_ascii=False),
|
||||
})
|
||||
|
||||
# 清除 ref_image_url,避免后续轮次重复附加
|
||||
ref_image_url = None
|
||||
# ref_image_url 不清除:消息构建只在循环外执行一次,
|
||||
# 后续工具调用仍需参考图(InstantStyle 等模型必须有 style_image)
|
||||
|
||||
yield {"type": "done", "data": {}}
|
||||
|
||||
@@ -38,20 +38,27 @@ async def execute_tool(
|
||||
tool_name: str,
|
||||
arguments: dict,
|
||||
ref_image_url: str | None = None,
|
||||
image_model: str | None = None,
|
||||
) -> dict:
|
||||
"""执行工具调用,返回结果。"""
|
||||
if tool_name == "generate_image":
|
||||
prompt = arguments["prompt"]
|
||||
num_images = arguments.get("num_images", 1)
|
||||
image_urls = await generate_images(
|
||||
result = await generate_images(
|
||||
prompt=prompt,
|
||||
num_images=num_images,
|
||||
ref_image_url=ref_image_url,
|
||||
model_id=image_model,
|
||||
)
|
||||
valid_urls = [u for u in result.urls if not u.startswith("[")]
|
||||
errors = [u for u in result.urls if u.startswith("[")]
|
||||
return {
|
||||
"success": True,
|
||||
"images": image_urls,
|
||||
"success": len(valid_urls) > 0,
|
||||
"images": valid_urls,
|
||||
"errors": errors,
|
||||
"prompt_used": prompt,
|
||||
"model_name": result.model_name,
|
||||
"model_id": result.model_id,
|
||||
}
|
||||
|
||||
return {"success": False, "error": f"未知工具: {tool_name}"}
|
||||
|
||||
Binary file not shown.
Binary file not shown.
@@ -7,6 +7,7 @@ from fastapi import APIRouter, File, Form, UploadFile
|
||||
from sse_starlette.sse import EventSourceResponse
|
||||
|
||||
from app.agent.loop import run_agent_loop
|
||||
from app.config import get_image_models_list, get_default_image_model_id
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -23,26 +24,56 @@ async def _save_upload(file: UploadFile) -> str:
|
||||
return f"/uploads/{filename}"
|
||||
|
||||
|
||||
@router.post("/upload-ref-image")
|
||||
async def upload_ref_image(
|
||||
file: UploadFile = File(...),
|
||||
):
|
||||
"""
|
||||
独立的参考图上传端点。
|
||||
前端选图后立即调用,返回服务端路径,供后续发消息时引用。
|
||||
"""
|
||||
UPLOADS_DIR.mkdir(parents=True, exist_ok=True)
|
||||
url = await _save_upload(file)
|
||||
return {"url": url, "filename": file.filename}
|
||||
|
||||
|
||||
@router.get("/models")
|
||||
async def list_models():
|
||||
"""返回可用的图像生成模型列表。"""
|
||||
return {
|
||||
"models": get_image_models_list(),
|
||||
"default": get_default_image_model_id(),
|
||||
}
|
||||
|
||||
|
||||
@router.post("/chat")
|
||||
async def chat(
|
||||
messages: str = Form(...),
|
||||
ref_image: Optional[UploadFile] = File(None),
|
||||
ref_image_url: Optional[str] = Form(None),
|
||||
image_model: Optional[str] = Form(None),
|
||||
):
|
||||
"""
|
||||
主对话端点。
|
||||
|
||||
参数:
|
||||
- messages: JSON 字符串,对话历史 [{role, content}]
|
||||
- ref_image: 可选的参考图文件
|
||||
- ref_image: 可选的参考图文件(兼容旧方式)
|
||||
- ref_image_url: 可选,已通过 /upload-ref-image 上传后的服务端路径
|
||||
- image_model: 可选,指定本次使用的生图模型短 ID
|
||||
"""
|
||||
parsed_messages = json.loads(messages)
|
||||
|
||||
ref_image_url: Optional[str] = None
|
||||
if ref_image and ref_image.filename:
|
||||
ref_image_url = await _save_upload(ref_image)
|
||||
resolved_ref_url: Optional[str] = None
|
||||
if ref_image_url:
|
||||
resolved_ref_url = ref_image_url
|
||||
elif ref_image and ref_image.filename:
|
||||
resolved_ref_url = await _save_upload(ref_image)
|
||||
|
||||
async def event_generator():
|
||||
async for event in run_agent_loop(parsed_messages, ref_image_url):
|
||||
async for event in run_agent_loop(
|
||||
parsed_messages, resolved_ref_url, image_model=image_model
|
||||
):
|
||||
yield {
|
||||
"event": event["type"],
|
||||
"data": json.dumps(event["data"], ensure_ascii=False),
|
||||
|
||||
156
art-agent/backend/app/config.py
Normal file
156
art-agent/backend/app/config.py
Normal file
@@ -0,0 +1,156 @@
|
||||
"""
|
||||
EPEEKit 集中配置。
|
||||
所有可调参数从环境变量读取,每次调用实时读取(不缓存),确保 load_dotenv() 后生效。
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
|
||||
def get_llm_model() -> str:
|
||||
return os.getenv("LLM_MODEL", "gpt-4o-mini")
|
||||
|
||||
|
||||
def get_llm_max_iterations() -> int:
|
||||
return int(os.getenv("LLM_MAX_ITERATIONS", "5"))
|
||||
|
||||
|
||||
# ─── 图像模型注册表 ─────────────────────────────────────
|
||||
#
|
||||
# 每个模型的配置说明:
|
||||
# id — 前端/API 使用的短 ID
|
||||
# name — 显示名称
|
||||
# provider — 生成服务提供者(对应 image_gen.py 中的 Provider)
|
||||
# model_id — Replicate 上的完整模型 ID
|
||||
# description — 前端下拉列表中的说明文字
|
||||
# supports_ref_image — 是否原生支持参考图输入(IP-Adapter 等)
|
||||
# ref_image_param — 传给 Replicate 的参考图参数名(模型间可能不同)
|
||||
# num_images_param — 批量生成参数名(Flux 用 num_outputs,Kolors 用 number_of_images)
|
||||
# default_params — 默认推理参数
|
||||
|
||||
IMAGE_MODELS: dict[str, dict[str, Any]] = {
|
||||
"flux-schnell": {
|
||||
"id": "flux-schnell",
|
||||
"name": "Flux Schnell",
|
||||
"provider": "replicate",
|
||||
"model_id": "black-forest-labs/flux-schnell",
|
||||
"description": "快速生成,适合快速迭代",
|
||||
"supports_ref_image": False,
|
||||
"num_images_param": "num_outputs",
|
||||
"default_params": {
|
||||
"aspect_ratio": "1:1",
|
||||
"output_format": "png",
|
||||
},
|
||||
},
|
||||
"flux-dev": {
|
||||
"id": "flux-dev",
|
||||
"name": "Flux Dev",
|
||||
"provider": "replicate",
|
||||
"model_id": "black-forest-labs/flux-dev",
|
||||
"description": "高质量生成,细节更好",
|
||||
"supports_ref_image": False,
|
||||
"num_images_param": "num_outputs",
|
||||
"default_params": {
|
||||
"aspect_ratio": "1:1",
|
||||
"output_format": "png",
|
||||
},
|
||||
},
|
||||
"sdxl": {
|
||||
"id": "sdxl",
|
||||
"name": "Stable Diffusion XL",
|
||||
"provider": "replicate",
|
||||
"model_id": "stability-ai/sdxl",
|
||||
"description": "经典 SDXL,支持 negative prompt",
|
||||
"supports_ref_image": False,
|
||||
"num_images_param": "num_outputs",
|
||||
"default_params": {
|
||||
"width": 1024,
|
||||
"height": 1024,
|
||||
"num_inference_steps": 50,
|
||||
"guidance_scale": 7.5,
|
||||
},
|
||||
},
|
||||
"instant-style": {
|
||||
"id": "instant-style",
|
||||
"name": "InstantStyle",
|
||||
"provider": "replicate",
|
||||
"model_id": "jyoung105/instant-style:c6f01e12f31cb99f9ee774a78992a71294f630a6f433d9aecfdc33b816fc4baa",
|
||||
"description": "强风格迁移,画风还原度高(较慢)",
|
||||
"supports_ref_image": True,
|
||||
"ref_image_param": "style_image",
|
||||
"num_images_param": "num_outputs",
|
||||
"default_params": {
|
||||
"width": 1024,
|
||||
"height": 1024,
|
||||
"num_inference_steps": 30,
|
||||
"guidance_scale": 5,
|
||||
"style_strength": 1.0,
|
||||
"block_mode": "style-only",
|
||||
"adapter_mode": "original",
|
||||
},
|
||||
},
|
||||
"kolors-ipadapter": {
|
||||
"id": "kolors-ipadapter",
|
||||
"name": "Kolors IP-Adapter",
|
||||
"provider": "replicate",
|
||||
"model_id": "fofr/kolors-with-ipadapter:5a1a92b2c0f81813225d48ed8e411813da41aa84e7582fb705d1af46eea36eed",
|
||||
"description": "风格参考生成,上传参考图效果最佳",
|
||||
"supports_ref_image": True,
|
||||
"ref_image_param": "image",
|
||||
"num_images_param": "number_of_images",
|
||||
"default_params": {
|
||||
"width": 1024,
|
||||
"height": 1024,
|
||||
"steps": 25,
|
||||
"cfg": 4,
|
||||
"ip_adapter_weight": 0.8,
|
||||
"ip_adapter_weight_type": "style transfer precise",
|
||||
"output_format": "png",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def get_default_image_model_id() -> str:
|
||||
"""返回 .env 中配置的默认模型短 ID,若不在注册表中则回退到 flux-schnell。"""
|
||||
env_model = os.getenv("IMAGE_MODEL", "flux-schnell")
|
||||
for mid, cfg in IMAGE_MODELS.items():
|
||||
if cfg["model_id"] == env_model or mid == env_model:
|
||||
return mid
|
||||
return "flux-schnell"
|
||||
|
||||
|
||||
def get_image_model_config(model_id: str | None = None) -> dict[str, Any]:
|
||||
"""根据短 ID 获取模型配置,未指定或不存在则使用默认模型。"""
|
||||
if model_id and model_id in IMAGE_MODELS:
|
||||
return IMAGE_MODELS[model_id]
|
||||
return IMAGE_MODELS[get_default_image_model_id()]
|
||||
|
||||
|
||||
def get_ref_image_model_id() -> str | None:
|
||||
"""返回有参考图时推荐使用的模型 ID(第一个 supports_ref_image=True 的模型)。"""
|
||||
for mid, cfg in IMAGE_MODELS.items():
|
||||
if cfg.get("supports_ref_image"):
|
||||
return mid
|
||||
return None
|
||||
|
||||
|
||||
def get_image_models_list() -> list[dict]:
|
||||
"""返回前端下拉列表所需的模型摘要信息。"""
|
||||
return [
|
||||
{
|
||||
"id": cfg["id"],
|
||||
"name": cfg["name"],
|
||||
"description": cfg["description"],
|
||||
"supports_ref_image": cfg.get("supports_ref_image", False),
|
||||
}
|
||||
for cfg in IMAGE_MODELS.values()
|
||||
]
|
||||
|
||||
|
||||
def get_image_aspect_ratio() -> str:
|
||||
return os.getenv("IMAGE_ASPECT_RATIO", "1:1")
|
||||
|
||||
|
||||
def get_image_output_format() -> str:
|
||||
return os.getenv("IMAGE_OUTPUT_FORMAT", "png")
|
||||
@@ -2,15 +2,15 @@ import os
|
||||
from pathlib import Path
|
||||
|
||||
from dotenv import load_dotenv
|
||||
load_dotenv(override=True)
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
|
||||
from app.api.chat import router as chat_router
|
||||
|
||||
load_dotenv()
|
||||
|
||||
app = FastAPI(title="Art Agent MVP")
|
||||
app = FastAPI(title="EPEEKit API")
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
|
||||
Binary file not shown.
Binary file not shown.
@@ -1,76 +1,220 @@
|
||||
"""Replicate 图像生成服务封装。"""
|
||||
"""图像生成服务 — Provider 抽象层。"""
|
||||
|
||||
import base64
|
||||
import mimetypes
|
||||
import os
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
import replicate
|
||||
from replicate import Client as ReplicateClient
|
||||
|
||||
GENERATED_DIR = Path(__file__).parent.parent.parent / "generated"
|
||||
from app.config import get_image_model_config, get_image_aspect_ratio, get_image_output_format
|
||||
|
||||
BACKEND_ROOT = Path(__file__).parent.parent.parent
|
||||
GENERATED_DIR = BACKEND_ROOT / "generated"
|
||||
|
||||
|
||||
def _make_replicate_client() -> ReplicateClient:
|
||||
"""
|
||||
创建 Replicate 客户端。
|
||||
Replicate SDK 内部创建 httpx transport 时不读取代理环境变量,
|
||||
需要我们手动把代理配置注入到 transport 中。
|
||||
"""
|
||||
timeout = httpx.Timeout(5.0, read=300.0, write=30.0, connect=30.0, pool=10.0)
|
||||
proxy_url = os.environ.get("HTTPS_PROXY") or os.environ.get("HTTP_PROXY")
|
||||
|
||||
transport_kwargs: dict[str, Any] = {}
|
||||
if proxy_url:
|
||||
transport_kwargs["proxy"] = proxy_url
|
||||
|
||||
our_transport = httpx.AsyncHTTPTransport(**transport_kwargs)
|
||||
|
||||
client = ReplicateClient(
|
||||
timeout=timeout,
|
||||
transport=our_transport,
|
||||
)
|
||||
|
||||
return client
|
||||
|
||||
|
||||
_replicate_client = _make_replicate_client()
|
||||
|
||||
|
||||
# ─── 公共工具 ──────────────────────────────────────────
|
||||
|
||||
def to_data_uri(image_path: str) -> str:
|
||||
"""将本地路径或已有 URL 转为 Replicate 可接受的格式(data URI 或原始 URL)。"""
|
||||
if image_path.startswith("data:") or image_path.startswith("http"):
|
||||
return image_path
|
||||
local = BACKEND_ROOT / image_path.lstrip("/")
|
||||
if local.exists():
|
||||
mime = mimetypes.guess_type(str(local))[0] or "image/png"
|
||||
b64 = base64.b64encode(local.read_bytes()).decode()
|
||||
return f"data:{mime};base64,{b64}"
|
||||
return image_path
|
||||
|
||||
|
||||
async def _download_image(url: str) -> str:
|
||||
"""下载远程图片到本地 generated/ 目录,返回本地 URL 路径。"""
|
||||
filename = f"{uuid.uuid4().hex}.png"
|
||||
filepath = GENERATED_DIR / filename
|
||||
async with httpx.AsyncClient() as client:
|
||||
proxy = os.environ.get("HTTPS_PROXY") or os.environ.get("HTTP_PROXY")
|
||||
async with httpx.AsyncClient(proxy=proxy, timeout=httpx.Timeout(60.0)) as client:
|
||||
resp = await client.get(url, follow_redirects=True)
|
||||
resp.raise_for_status()
|
||||
filepath.write_bytes(resp.content)
|
||||
return f"/generated/{filename}"
|
||||
|
||||
|
||||
# ─── Provider 抽象基类 ─────────────────────────────────
|
||||
|
||||
class ImageProvider(ABC):
|
||||
"""所有图像生成 provider 的基类。"""
|
||||
|
||||
@abstractmethod
|
||||
async def generate(
|
||||
self,
|
||||
model_config: dict[str, Any],
|
||||
prompt: str,
|
||||
num_images: int = 1,
|
||||
ref_image_url: str | None = None,
|
||||
) -> list[str]:
|
||||
"""生成图片并返回本地 URL 列表。"""
|
||||
...
|
||||
|
||||
|
||||
class ReplicateProvider(ImageProvider):
|
||||
"""通过 Replicate API 调用模型。"""
|
||||
|
||||
async def generate(
|
||||
self,
|
||||
model_config: dict[str, Any],
|
||||
prompt: str,
|
||||
num_images: int = 1,
|
||||
ref_image_url: str | None = None,
|
||||
) -> list[str]:
|
||||
replicate_model_id = model_config["model_id"]
|
||||
default_params: dict[str, Any] = model_config.get("default_params", {})
|
||||
|
||||
local_urls: list[str] = []
|
||||
|
||||
try:
|
||||
input_params = self._build_input(
|
||||
model_config, default_params, prompt, num_images, ref_image_url
|
||||
)
|
||||
|
||||
output = await _replicate_client.async_run(
|
||||
replicate_model_id, input=input_params, wait=False
|
||||
)
|
||||
|
||||
items = output if isinstance(output, list) else [output]
|
||||
for item in items:
|
||||
url = str(item)
|
||||
if url.startswith("https://") or url.startswith("http://") or url.startswith("data:"):
|
||||
local_urls.append(await _download_image(url))
|
||||
else:
|
||||
local_urls.append(f"[生成失败: 模型返回非图片内容: {url[:200]}]")
|
||||
|
||||
except Exception as e:
|
||||
detail = str(e) or f"{type(e).__name__}: {repr(e)}"
|
||||
local_urls.append(f"[生成失败: {detail}]")
|
||||
|
||||
return local_urls
|
||||
|
||||
@staticmethod
|
||||
def _build_input(
|
||||
model_config: dict[str, Any],
|
||||
default_params: dict[str, Any],
|
||||
prompt: str,
|
||||
num_images: int = 1,
|
||||
ref_image_url: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
根据模型配置构建 Replicate input 参数。
|
||||
|
||||
通过 model_config 中的标志字段自动适配不同模型:
|
||||
- supports_ref_image / ref_image_param: 参考图注入
|
||||
- num_images_param: 各模型的批量生成参数名(如 number_of_images / num_outputs)
|
||||
"""
|
||||
model_id = model_config["model_id"]
|
||||
supports_ref = model_config.get("supports_ref_image", False)
|
||||
ref_param_name = model_config.get("ref_image_param", "image")
|
||||
num_images_param = model_config.get("num_images_param", "num_outputs")
|
||||
|
||||
# ── 支持参考图的模型(IP-Adapter 系列)──
|
||||
if supports_ref:
|
||||
params: dict[str, Any] = {"prompt": prompt}
|
||||
for k, v in default_params.items():
|
||||
params[k] = v
|
||||
params[num_images_param] = num_images
|
||||
if ref_image_url:
|
||||
params[ref_param_name] = to_data_uri(ref_image_url)
|
||||
return params
|
||||
|
||||
# ── Flux 系列(纯文生图)──
|
||||
is_flux = "flux" in model_id.lower()
|
||||
params = {"prompt": prompt, num_images_param: num_images}
|
||||
|
||||
if is_flux:
|
||||
params["aspect_ratio"] = default_params.get(
|
||||
"aspect_ratio", get_image_aspect_ratio()
|
||||
)
|
||||
params["output_format"] = default_params.get(
|
||||
"output_format", get_image_output_format()
|
||||
)
|
||||
else:
|
||||
# SDXL 类模型
|
||||
params["width"] = default_params.get("width", 1024)
|
||||
params["height"] = default_params.get("height", 1024)
|
||||
if "num_inference_steps" in default_params:
|
||||
params["num_inference_steps"] = default_params["num_inference_steps"]
|
||||
if "guidance_scale" in default_params:
|
||||
params["guidance_scale"] = default_params["guidance_scale"]
|
||||
|
||||
return params
|
||||
|
||||
|
||||
# ─── Provider 注册 ─────────────────────────────────────
|
||||
|
||||
_PROVIDERS: dict[str, ImageProvider] = {
|
||||
"replicate": ReplicateProvider(),
|
||||
}
|
||||
|
||||
|
||||
# ─── 公开 API ──────────────────────────────────────────
|
||||
|
||||
class GenerateResult:
|
||||
"""图片生成结果,包含生成的 URL 列表和实际使用的模型信息。"""
|
||||
|
||||
def __init__(self, urls: list[str], model_name: str, model_id: str):
|
||||
self.urls = urls
|
||||
self.model_name = model_name
|
||||
self.model_id = model_id
|
||||
|
||||
|
||||
async def generate_images(
|
||||
prompt: str,
|
||||
num_images: int = 1,
|
||||
ref_image_url: str | None = None,
|
||||
) -> list[str]:
|
||||
model_id: str | None = None,
|
||||
) -> GenerateResult:
|
||||
"""
|
||||
调用 Replicate 生成图片。
|
||||
统一入口:根据 model_id 查注册表,分发到对应 provider。
|
||||
|
||||
返回本地可访问的图片 URL 列表。
|
||||
model_id 为空时使用 .env 中配置的默认模型。
|
||||
"""
|
||||
local_urls = []
|
||||
config = get_image_model_config(model_id)
|
||||
provider_name = config.get("provider", "replicate")
|
||||
provider = _PROVIDERS.get(provider_name)
|
||||
|
||||
for _ in range(num_images):
|
||||
try:
|
||||
if ref_image_url and ref_image_url.startswith("/"):
|
||||
# 本地路径转为 file URI 不适用于 Replicate,
|
||||
# 需要用户上传的图先通过后端 URL 访问
|
||||
# MVP 阶段:参考图作为 prompt 的文字补充,不直接传给模型
|
||||
# 仅使用 flux-schnell 文生图
|
||||
output = await replicate.async_run(
|
||||
"black-forest-labs/flux-schnell",
|
||||
input={
|
||||
"prompt": prompt,
|
||||
"num_outputs": 1,
|
||||
"aspect_ratio": "1:1",
|
||||
"output_format": "png",
|
||||
},
|
||||
)
|
||||
else:
|
||||
output = await replicate.async_run(
|
||||
"black-forest-labs/flux-schnell",
|
||||
input={
|
||||
"prompt": prompt,
|
||||
"num_outputs": 1,
|
||||
"aspect_ratio": "1:1",
|
||||
"output_format": "png",
|
||||
},
|
||||
)
|
||||
model_name = config.get("name", "未知模型")
|
||||
resolved_id = config.get("id", model_id or "unknown")
|
||||
|
||||
# output 是 FileOutput 列表或单个 URL
|
||||
if isinstance(output, list):
|
||||
for item in output:
|
||||
url = str(item)
|
||||
local_url = await _download_image(url)
|
||||
local_urls.append(local_url)
|
||||
else:
|
||||
url = str(output)
|
||||
local_url = await _download_image(url)
|
||||
local_urls.append(local_url)
|
||||
if not provider:
|
||||
return GenerateResult([f"[未知 provider: {provider_name}]"], model_name, resolved_id)
|
||||
|
||||
except Exception as e:
|
||||
local_urls.append(f"[生成失败: {e}]")
|
||||
|
||||
return local_urls
|
||||
urls = await provider.generate(config, prompt, num_images, ref_image_url)
|
||||
return GenerateResult(urls, model_name, resolved_id)
|
||||
|
||||
Reference in New Issue
Block a user