Files
EPEEAIKit/art-agent/backend/app/services/view_transform.py
2026-04-20 21:52:35 +08:00

358 lines
13 KiB
Python
Raw Permalink 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.
"""视角变换服务 — 调度单图→多视角生成,支持两种管道:
- grid 管道Zero123++ 直接输出 6 宫格拼接图,切分后得到固定 6 视角。
- mesh 管道Trellis/Hunyuan3D → .glb + 预渲染视频 → 抽帧 → 可选风格还原。
"""
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
from app.services import mesh_generator, style_restorer, video_frame_extractor
DEFAULT_MESH_AZIMUTHS = [0, 60, 120, 180, 240, 300]
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 中的视角。
由于不同 Replicate 部署版本输出的拼接图布局可能不同(有的是 3列×2行 的
宽图,有的是 2列×3行 的高图),这里会根据实际图像长宽比自动纠正 cols/rows——
假定每个单视角为正方形,按 aspect ratio 推断真实布局,防止切图错位。
"""
img = Image.open(image_path)
w, h = img.size
expected_cells = cols * rows
if expected_cells > 0 and len(view_labels) == expected_cells:
# 假定每个视角为正方形:真实 aspect ratio = cols / rows
configured_ratio = cols / rows
actual_ratio = w / h if h else configured_ratio
# 如果实际 ratio 与配置差异较大(超过 20%),判定为布局反了
if actual_ratio > 0 and (
max(configured_ratio, actual_ratio) / min(configured_ratio, actual_ratio) > 1.2
):
logger.warning(
"Grid 布局与实际图像不匹配:配置 cols=%d rows=%dratio=%.2f"
"实际图像 %dx%dratio=%.2f),自动交换 cols/rows。",
cols, rows, configured_ratio, w, h, actual_ratio,
)
cols, rows = rows, cols
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,
azimuths: list[int] | None = None,
elevations: list[int] | None = None,
preserve_style: bool = True,
) -> dict[str, Any]:
"""调用视角变换模型,返回多视角图片列表。
Args:
image_path: 输入图片路径本地路径、data URI 或 URL
model_id: 视角变换模型 IDzero123plus / trellis / hunyuan3d
azimuths: mesh 管道专用,方位角数组
elevations: mesh 管道专用,仰角数组(当前一期仅 0 有效)
preserve_style: mesh 管道专用,是否走阶段 3 风格还原
Returns:
统一字段(不同管道可能缺省):
{
"success": bool,
"images": [{"url": str, "azimuth": int, "elevation": int}, ...],
"grid_image": str | None, # grid 管道:原始拼接图 URL
"mesh_url": str | None, # mesh 管道:.glb URL
"color_video": str | None, # mesh 管道360° color video URL
"normal_video": str | None, # mesh 管道:法线 video URL
"pipeline": "grid" | "mesh",
"model_name": str,
"error": str | None,
}
"""
config = get_view_transform_model_config(model_id)
pipeline = config.get("pipeline", "grid")
GENERATED_DIR.mkdir(parents=True, exist_ok=True)
if pipeline == "mesh":
return await _transform_view_mesh(
image_path=image_path,
config=config,
azimuths=azimuths,
elevations=elevations,
preserve_style=preserve_style,
)
return await _transform_view_grid(image_path=image_path, config=config)
async def _transform_view_grid(
image_path: str,
config: dict[str, Any],
) -> dict[str, Any]:
"""原 Zero123++ 逻辑:单次调用 → grid 图 → 切 6 张。"""
model_name = config["name"]
replicate_model_id = config["model_id"]
grid_layout = config["grid_layout"]
output_views = config["output_views"]
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,
"mesh_url": None,
"color_video": None,
"normal_video": None,
"pipeline": "grid",
"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,
"mesh_url": None,
"color_video": None,
"normal_video": None,
"pipeline": "grid",
"model_name": model_name,
"error": detail,
}
async def _transform_view_mesh(
image_path: str,
config: dict[str, Any],
azimuths: list[int] | None,
elevations: list[int] | None,
preserve_style: bool,
) -> dict[str, Any]:
"""Mesh 管道调度器:阶段 1 生成 mesh + video阶段 2/3 后续 Phase 接入。
Phase 1 验收状态:只跑阶段 1返回 .glb / color_video / normal_video 本地 URL
images 列表为空。Phase 2 完成后会填充 imagesraw 抽帧结果),
Phase 3 完成后 images 会是经 IP-Adapter+ControlNet 重绘的最终图。
"""
model_name = config["name"]
model_id = config["id"]
# === Stage 1生成 mesh + 预渲染视频 ===
stage1 = await mesh_generator.generate_mesh(image_path, model_id=model_id)
if not stage1.get("success"):
return {
"success": False,
"images": [],
"grid_image": None,
"mesh_url": None,
"color_video": None,
"normal_video": None,
"combined_video": None,
"pipeline": "mesh",
"model_name": model_name,
"error": stage1.get("error") or "阶段 1 失败",
}
# === Stage 2抽帧 + 切左右半color + normal===
# azimuths / elevations 缺省时用默认 6 视角;长度对齐
if not azimuths:
azimuths = list(DEFAULT_MESH_AZIMUTHS)
if elevations is None or not elevations:
elevations = [0] * len(azimuths)
elif len(elevations) < len(azimuths):
elevations = list(elevations) + [0] * (len(azimuths) - len(elevations))
elif len(elevations) > len(azimuths):
elevations = list(elevations[: len(azimuths)])
# 优先 color_video当前 Trellis 部署版恒为 None缺则 fallback 到 combined_video
combined_video_path = stage1.get("combined_video_path")
frames: list[dict[str, Any]] = []
stage2_error: str | None = None
if combined_video_path:
try:
frames = await video_frame_extractor.extract_by_azimuth(
combined_video_path=combined_video_path,
azimuths=azimuths,
elevations=elevations,
total_rotation_deg=360,
)
except Exception as e:
stage2_error = f"Stage 2 抽帧失败:{e}"
logger.error(stage2_error, exc_info=True)
else:
stage2_error = "Trellis 未返回 combined_video无法执行 Stage 2 抽帧"
logger.error(stage2_error)
# === Stage 3风格还原可选===
# preserve_style=True 时对每一帧的 normal_png 做 IP-Adapter+ControlNet 重绘,
# 用户原图做风格参考;失败的单帧会回退为 Stage 2 的 color 帧,不阻塞其它帧。
stage3_error: str | None = None
stage3_results: list[dict[str, Any]] = []
did_stage3 = bool(preserve_style) and bool(frames)
if did_stage3:
try:
stage3_results = await style_restorer.restore_batch(
original_image=image_path,
frames=frames,
concurrency=3,
)
except Exception as e:
stage3_error = f"Stage 3 风格还原失败:{e}"
logger.error(stage3_error, exc_info=True)
# 组装最终 images
# - Stage 3 成功的帧 → 用重绘后的 URL
# - Stage 3 失败 / 未跑 → 回退到 Stage 2 的 color 帧(保证前端至少能看到点东西)
images: list[dict[str, Any]] = []
for i, f in enumerate(frames):
restyled = stage3_results[i] if i < len(stage3_results) else None
if restyled and restyled.get("success") and restyled.get("url"):
final_url = restyled["url"]
frame_error = None
else:
final_url = f["color_png"]
frame_error = (restyled or {}).get("error") if did_stage3 else None
images.append({
"url": final_url,
"azimuth": f["az"],
"elevation": f["el"],
"normal_url": f["normal_png"],
"color_url": f["color_png"],
"frame_idx": f["frame_idx"],
"restyled": bool(restyled and restyled.get("success")),
**({"restyle_error": frame_error} if frame_error else {}),
})
# Stage 2 失败不把整体置为失败——Stage 1 的 mesh / video 产物仍有交付价值
# Stage 3 部分失败也不置为整体失败(每帧都有 Stage 2 回退)
success = bool(stage1.get("success")) and (stage2_error is None)
# 合并错误信息Stage 2 和 Stage 3 错误都要暴露给前端)
error_parts = [e for e in (stage2_error, stage3_error) if e]
error_out = "; ".join(error_parts) if error_parts else None
if did_stage3:
restyled_ok = sum(1 for img in images if img.get("restyled"))
stage_label = f"stage3 ({restyled_ok}/{len(images)} restyled)"
elif frames:
stage_label = "stage2"
else:
stage_label = "stage1"
return {
"success": success,
"images": images,
"grid_image": None,
"mesh_url": stage1.get("glb_path"),
"color_video": stage1.get("color_video_path"),
"normal_video": stage1.get("normal_video_path"),
"combined_video": combined_video_path,
"pipeline": "mesh",
"model_name": model_name,
"error": error_out,
"_stage": stage_label,
"_preserve_style_requested": bool(preserve_style),
}