引入3d模型
This commit is contained in:
@@ -1,4 +1,8 @@
|
||||
"""视角变换服务 — 调用 Zero123++ 等模型,将单张图片转为多视角图片。"""
|
||||
"""视角变换服务 — 调度单图→多视角生成,支持两种管道:
|
||||
|
||||
- grid 管道:Zero123++ 直接输出 6 宫格拼接图,切分后得到固定 6 视角。
|
||||
- mesh 管道:Trellis/Hunyuan3D → .glb + 预渲染视频 → 抽帧 → 可选风格还原。
|
||||
"""
|
||||
|
||||
import base64
|
||||
import logging
|
||||
@@ -13,6 +17,9 @@ 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__)
|
||||
|
||||
@@ -29,9 +36,30 @@ def _split_grid_image(
|
||||
|
||||
返回格式:[{"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=%d(ratio=%.2f),"
|
||||
"实际图像 %dx%d(ratio=%.2f),自动交换 cols/rows。",
|
||||
cols, rows, configured_ratio, w, h, actual_ratio,
|
||||
)
|
||||
cols, rows = rows, cols
|
||||
|
||||
cell_w = w // cols
|
||||
cell_h = h // rows
|
||||
|
||||
@@ -73,30 +101,60 @@ async def _download_to_local(url: str) -> Path:
|
||||
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: 视角变换模型 ID,默认 zero123plus
|
||||
model_id: 视角变换模型 ID(zero123plus / trellis / hunyuan3d)
|
||||
azimuths: mesh 管道专用,方位角数组
|
||||
elevations: mesh 管道专用,仰角数组(当前一期仅 0 有效)
|
||||
preserve_style: mesh 管道专用,是否走阶段 3 风格还原
|
||||
|
||||
Returns:
|
||||
统一字段(不同管道可能缺省):
|
||||
{
|
||||
"success": bool,
|
||||
"images": [{"url": str, "azimuth": int, "elevation": int}, ...],
|
||||
"grid_image": str, # 原始拼接图的本地 URL
|
||||
"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"]
|
||||
|
||||
GENERATED_DIR.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
try:
|
||||
image_uri = to_data_uri(image_path)
|
||||
input_params = {"image": image_uri}
|
||||
@@ -141,6 +199,10 @@ async def transform_view(
|
||||
"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,
|
||||
}
|
||||
@@ -152,6 +214,144 @@ async def transform_view(
|
||||
"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 完成后会填充 images(raw 抽帧结果),
|
||||
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),
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user