Files
EPEEAIKit/art-agent/backend/app/services/view_transform.py

158 lines
4.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""视角变换服务 — 调用 Zero123++ 等模型,将单张图片转为多视角图片。"""
import base64
import logging
import mimetypes
import os
import uuid
from pathlib import Path
from typing import Any
import httpx
from PIL import Image
from app.config import get_view_transform_model_config
from app.services.image_gen import _make_replicate_client, GENERATED_DIR, to_data_uri
logger = logging.getLogger(__name__)
_replicate_client = _make_replicate_client()
def _split_grid_image(
image_path: Path,
cols: int,
rows: int,
view_labels: list[dict[str, int]],
) -> list[dict[str, Any]]:
"""将 Zero123++ 输出的 grid 拼接图切分为独立的视角图片。
返回格式:[{"url": "/generated/xxx.png", "azimuth": 30, "elevation": 30}, ...]
Grid 布局从左到右、从上到下依次对应 view_labels 中的视角。
"""
img = Image.open(image_path)
w, h = img.size
cell_w = w // cols
cell_h = h // rows
results: list[dict[str, Any]] = []
for idx, view in enumerate(view_labels):
row_idx = idx // cols
col_idx = idx % cols
left = col_idx * cell_w
upper = row_idx * cell_h
right = left + cell_w
lower = upper + cell_h
cropped = img.crop((left, upper, right, lower))
filename = f"view_{uuid.uuid4().hex[:8]}_az{view['azimuth']}_el{view['elevation']}.png"
filepath = GENERATED_DIR / filename
cropped.save(filepath, "PNG")
results.append({
"url": f"/generated/{filename}",
"azimuth": view["azimuth"],
"elevation": view["elevation"],
})
return results
async def _download_to_local(url: str) -> Path:
"""下载远程图片到本地 generated/ 目录,返回本地路径。"""
filename = f"grid_{uuid.uuid4().hex}.png"
filepath = GENERATED_DIR / filename
proxy = os.environ.get("HTTPS_PROXY") or os.environ.get("HTTP_PROXY")
async with httpx.AsyncClient(proxy=proxy, timeout=httpx.Timeout(120.0)) as client:
resp = await client.get(url, follow_redirects=True)
resp.raise_for_status()
filepath.write_bytes(resp.content)
return filepath
async def transform_view(
image_path: str,
model_id: str | None = None,
) -> dict[str, Any]:
"""调用视角变换模型,返回多视角图片列表。
Args:
image_path: 输入图片路径本地路径、data URI 或 URL
model_id: 视角变换模型 ID默认 zero123plus
Returns:
{
"success": bool,
"images": [{"url": str, "azimuth": int, "elevation": int}, ...],
"grid_image": str, # 原始拼接图的本地 URL
"model_name": str,
"error": str | None,
}
"""
config = get_view_transform_model_config(model_id)
model_name = config["name"]
replicate_model_id = config["model_id"]
grid_layout = config["grid_layout"]
output_views = config["output_views"]
GENERATED_DIR.mkdir(parents=True, exist_ok=True)
try:
image_uri = to_data_uri(image_path)
input_params = {"image": image_uri}
logger.info("调用 %s 进行视角变换...", model_name)
output = await _replicate_client.async_run(
replicate_model_id, input=input_params, wait=False
)
# Zero123++ 返回单张 grid 拼接图
items = output if isinstance(output, list) else [output]
if not items:
return {
"success": False,
"images": [],
"grid_image": None,
"model_name": model_name,
"error": "模型未返回任何输出",
}
grid_url = str(items[0])
if not (grid_url.startswith("http://") or grid_url.startswith("https://")):
return {
"success": False,
"images": [],
"grid_image": None,
"model_name": model_name,
"error": f"模型返回非图片内容: {grid_url[:200]}",
}
grid_local_path = await _download_to_local(grid_url)
grid_local_url = f"/generated/{grid_local_path.name}"
view_images = _split_grid_image(
grid_local_path,
cols=grid_layout["cols"],
rows=grid_layout["rows"],
view_labels=output_views,
)
return {
"success": True,
"images": view_images,
"grid_image": grid_local_url,
"model_name": model_name,
"error": None,
}
except Exception as e:
detail = str(e) or f"{type(e).__name__}: {repr(e)}"
logger.error("视角变换失败: %s", detail, exc_info=True)
return {
"success": False,
"images": [],
"grid_image": None,
"model_name": model_name,
"error": detail,
}