313 lines
11 KiB
Python
313 lines
11 KiB
Python
"""
|
||
人脸贴图渲染核心:接受一张图 + 一个效果 ID,输出 PNG 字节。
|
||
|
||
流程(与 Android 端 vulkan/FaceApp.cpp 一致,CPU 实现):
|
||
1. MediaPipe FaceLandmarker → 478 normalized landmarks per face
|
||
2. 取最大人脸(按 landmark bbox 面积)
|
||
3. 把 obj 的 468 顶点位置改写为 landmarks[INDEX_MAP[i]],得到 dst 顶点
|
||
4. 对 mesh 的每个三角形:从 512×512 effect 贴图取 src 三角形,仿射 warp
|
||
到 dst 三角形位置,按贴图 alpha 通道 alpha-blend 到原图副本上
|
||
5. 输出 PNG(保持原图尺寸)
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import io
|
||
import threading
|
||
from dataclasses import dataclass
|
||
from pathlib import Path
|
||
|
||
import cv2
|
||
import numpy as np
|
||
from PIL import Image, ImageOps
|
||
|
||
from .index_map import INDEX_MAP
|
||
from .mesh import FaceMesh, load_face_mesh
|
||
|
||
|
||
class NoFaceError(Exception):
|
||
"""图里检测不到人脸。"""
|
||
|
||
|
||
class EffectNotFoundError(Exception):
|
||
"""指定 effect_id 没有对应贴图。"""
|
||
|
||
|
||
@dataclass
|
||
class RendererPaths:
|
||
obj_path: Path
|
||
effects_dir: Path
|
||
|
||
|
||
class FaceRenderer:
|
||
"""
|
||
单例式渲染器:mesh / mediapipe 模型加载一次复用。
|
||
|
||
线程安全:MediaPipe FaceLandmarker 的 detect() 不保证多线程并发,
|
||
用 _lock 串行化。FastAPI 单进程同步部署下,请求本身就串行,开销可忽略。
|
||
"""
|
||
|
||
def __init__(self, paths: RendererPaths, max_image_bytes: int = 20 * 1024 * 1024):
|
||
self.paths = paths
|
||
self.max_image_bytes = max_image_bytes
|
||
self._mesh: FaceMesh = load_face_mesh(paths.obj_path)
|
||
if len(self._mesh.uvs) != len(INDEX_MAP):
|
||
raise RuntimeError(
|
||
f"mesh has {len(self._mesh.uvs)} vertices but INDEX_MAP has "
|
||
f"{len(INDEX_MAP)}; obj/index_map.py 不匹配"
|
||
)
|
||
# 预转 numpy 加速
|
||
self._index_map = np.asarray(INDEX_MAP, dtype=np.int32)
|
||
|
||
self._effect_cache: dict[int, np.ndarray] = {}
|
||
self._cache_lock = threading.Lock()
|
||
|
||
# MediaPipe 延迟加载(首次请求时加载,缩短启动时间,便于本地调试)
|
||
self._landmarker = None
|
||
self._landmarker_lock = threading.Lock()
|
||
|
||
# ---------- public ----------
|
||
|
||
def list_effect_ids(self) -> list[int]:
|
||
"""枚举 effects/<n>.png 中的 n。"""
|
||
ids: list[int] = []
|
||
if not self.paths.effects_dir.is_dir():
|
||
return ids
|
||
for p in self.paths.effects_dir.iterdir():
|
||
if p.suffix.lower() != ".png":
|
||
continue
|
||
try:
|
||
n = int(p.stem)
|
||
except ValueError:
|
||
continue
|
||
if 1 <= n <= 99:
|
||
ids.append(n)
|
||
return sorted(ids)
|
||
|
||
def render(self, image_bytes: bytes, effect_id: int) -> bytes:
|
||
if len(image_bytes) > self.max_image_bytes:
|
||
raise ValueError("image too large")
|
||
|
||
effect_rgba = self._load_effect(effect_id) # (H,W,4) uint8
|
||
|
||
# 1. 解码 + 处理 EXIF orientation → RGB ndarray
|
||
img_rgb = self._decode_input(image_bytes)
|
||
h, w = img_rgb.shape[:2]
|
||
|
||
# 2. MediaPipe 检测,挑最大人脸
|
||
landmarks_norm = self._detect_largest_face(img_rgb) # (N,2) ∈ [0,1]² (N>=468)
|
||
|
||
# 3. obj 顶点 → 像素坐标
|
||
dst_xy = self._compute_dst_vertices(landmarks_norm, w, h)
|
||
|
||
# 4. warp + 混合
|
||
out = img_rgb.copy()
|
||
self._composite(out, effect_rgba, dst_xy)
|
||
|
||
# 5. 编码 PNG
|
||
return self._encode_png(out)
|
||
|
||
# ---------- internals ----------
|
||
|
||
def _load_effect(self, effect_id: int) -> np.ndarray:
|
||
if not (1 <= effect_id <= 99):
|
||
raise EffectNotFoundError(f"effect_id out of range: {effect_id}")
|
||
with self._cache_lock:
|
||
cached = self._effect_cache.get(effect_id)
|
||
if cached is not None:
|
||
return cached
|
||
png_path = self.paths.effects_dir / f"{effect_id}.png"
|
||
if not png_path.is_file():
|
||
raise EffectNotFoundError(f"effect {effect_id} not found")
|
||
# 用 PIL 解码以保证 RGBA 一致;OpenCV imread 对部分 PNG 会丢 alpha
|
||
img = Image.open(png_path).convert("RGBA")
|
||
arr = np.array(img, dtype=np.uint8) # (H,W,4) RGBA
|
||
with self._cache_lock:
|
||
self._effect_cache[effect_id] = arr
|
||
return arr
|
||
|
||
def _decode_input(self, image_bytes: bytes) -> np.ndarray:
|
||
try:
|
||
pil = Image.open(io.BytesIO(image_bytes))
|
||
pil = ImageOps.exif_transpose(pil) # 处理手机 EXIF 旋转
|
||
pil = pil.convert("RGB")
|
||
except Exception as e:
|
||
raise ValueError(f"invalid image: {e}") from e
|
||
return np.array(pil, dtype=np.uint8)
|
||
|
||
def _ensure_landmarker(self):
|
||
if self._landmarker is not None:
|
||
return
|
||
with self._landmarker_lock:
|
||
if self._landmarker is not None:
|
||
return
|
||
# 延迟 import,避免单元测试时强依赖 mediapipe
|
||
from mediapipe.tasks import python as mp_python
|
||
from mediapipe.tasks.python import vision as mp_vision
|
||
|
||
model_path = self._find_model_file()
|
||
base_options = mp_python.BaseOptions(model_asset_path=str(model_path))
|
||
options = mp_vision.FaceLandmarkerOptions(
|
||
base_options=base_options,
|
||
running_mode=mp_vision.RunningMode.IMAGE,
|
||
num_faces=5, # 检测最多 5 张,挑最大那张
|
||
min_face_detection_confidence=0.5,
|
||
min_face_presence_confidence=0.5,
|
||
min_tracking_confidence=0.5,
|
||
output_face_blendshapes=False,
|
||
output_facial_transformation_matrixes=False,
|
||
)
|
||
self._landmarker = mp_vision.FaceLandmarker.create_from_options(options)
|
||
|
||
def _find_model_file(self) -> Path:
|
||
"""优先用本服务 assets 下的,找不到再回退到仓库 app/src/main/assets。"""
|
||
local = self.paths.obj_path.parent / "face_landmarker.task"
|
||
if local.is_file():
|
||
return local
|
||
# 回退到仓库 app/src/main/assets/face_landmarker.task(Android SDK 用同一份)
|
||
repo_root = self.paths.obj_path.resolve().parents[3]
|
||
fallback = repo_root / "app" / "src" / "main" / "assets" / "face_landmarker.task"
|
||
if fallback.is_file():
|
||
return fallback
|
||
raise FileNotFoundError(
|
||
"face_landmarker.task not found; place it under "
|
||
f"{local} or {fallback}"
|
||
)
|
||
|
||
def _detect_largest_face(self, img_rgb: np.ndarray) -> np.ndarray:
|
||
self._ensure_landmarker()
|
||
import mediapipe as mp
|
||
|
||
mp_image = mp.Image(image_format=mp.ImageFormat.SRGB, data=img_rgb)
|
||
with self._landmarker_lock:
|
||
result = self._landmarker.detect(mp_image)
|
||
|
||
faces = result.face_landmarks or []
|
||
if not faces:
|
||
raise NoFaceError("no face detected")
|
||
|
||
# 选 bbox 面积最大的那张
|
||
def bbox_area(landmarks) -> float:
|
||
xs = [lm.x for lm in landmarks]
|
||
ys = [lm.y for lm in landmarks]
|
||
return (max(xs) - min(xs)) * (max(ys) - min(ys))
|
||
|
||
biggest = max(faces, key=bbox_area)
|
||
# 转 ndarray,只取 xy
|
||
n = len(biggest)
|
||
arr = np.empty((n, 2), dtype=np.float32)
|
||
for i, lm in enumerate(biggest):
|
||
arr[i, 0] = lm.x
|
||
arr[i, 1] = lm.y
|
||
return arr
|
||
|
||
def _compute_dst_vertices(
|
||
self, landmarks_norm: np.ndarray, w: int, h: int
|
||
) -> np.ndarray:
|
||
"""
|
||
landmarks_norm: (N, 2) ∈ [0,1]²
|
||
返回 (468, 2) float32 像素坐标,对应 mesh 的 468 个顶点。
|
||
"""
|
||
n_landmarks = landmarks_norm.shape[0]
|
||
if self._index_map.max() >= n_landmarks:
|
||
raise RuntimeError(
|
||
f"INDEX_MAP references landmark index {int(self._index_map.max())} "
|
||
f"but mediapipe returned only {n_landmarks} landmarks"
|
||
)
|
||
picked = landmarks_norm[self._index_map] # (468, 2)
|
||
out = np.empty_like(picked)
|
||
out[:, 0] = picked[:, 0] * w
|
||
out[:, 1] = picked[:, 1] * h
|
||
return out.astype(np.float32)
|
||
|
||
def _composite(
|
||
self,
|
||
canvas_rgb: np.ndarray,
|
||
effect_rgba: np.ndarray,
|
||
dst_xy: np.ndarray,
|
||
) -> None:
|
||
"""
|
||
在 canvas_rgb(uint8 H×W×3,原地修改)上对 mesh 每个三角形:
|
||
从 effect_rgba 取对应 src 三角形 → warpAffine 到 dst → alpha 混合。
|
||
"""
|
||
eh, ew = effect_rgba.shape[:2]
|
||
canvas_h, canvas_w = canvas_rgb.shape[:2]
|
||
uvs = self._mesh.uvs # (468, 2)
|
||
tris = self._mesh.triangles # (852, 3)
|
||
|
||
# 把 UV 一次性转成 effect 像素坐标
|
||
src_pts_all = np.empty_like(uvs)
|
||
src_pts_all[:, 0] = uvs[:, 0] * ew
|
||
src_pts_all[:, 1] = uvs[:, 1] * eh
|
||
|
||
effect_rgb = effect_rgba[..., :3]
|
||
effect_a = effect_rgba[..., 3]
|
||
|
||
for tri in tris:
|
||
i0, i1, i2 = int(tri[0]), int(tri[1]), int(tri[2])
|
||
src_tri = np.float32([src_pts_all[i0], src_pts_all[i1], src_pts_all[i2]])
|
||
dst_tri = np.float32([dst_xy[i0], dst_xy[i1], dst_xy[i2]])
|
||
|
||
# 用 dst bbox 限定计算区域,避免对全画布 warp
|
||
x_min = int(np.floor(dst_tri[:, 0].min()))
|
||
y_min = int(np.floor(dst_tri[:, 1].min()))
|
||
x_max = int(np.ceil(dst_tri[:, 0].max()))
|
||
y_max = int(np.ceil(dst_tri[:, 1].max()))
|
||
|
||
# clip 到画布
|
||
x_min_c = max(0, x_min)
|
||
y_min_c = max(0, y_min)
|
||
x_max_c = min(canvas_w, x_max)
|
||
y_max_c = min(canvas_h, y_max)
|
||
if x_max_c <= x_min_c or y_max_c <= y_min_c:
|
||
continue
|
||
|
||
box_w = x_max_c - x_min_c
|
||
box_h = y_max_c - y_min_c
|
||
|
||
# 构造从 src → 局部 bbox 坐标系(左上为原点)的仿射
|
||
dst_tri_local = dst_tri - np.array([x_min_c, y_min_c], dtype=np.float32)
|
||
M = cv2.getAffineTransform(src_tri, dst_tri_local)
|
||
|
||
# 给 RGB 和 alpha 分别 warp 到 bbox 大小
|
||
warped_rgb = cv2.warpAffine(
|
||
effect_rgb,
|
||
M,
|
||
(box_w, box_h),
|
||
flags=cv2.INTER_LINEAR,
|
||
borderMode=cv2.BORDER_REPLICATE,
|
||
)
|
||
warped_a = cv2.warpAffine(
|
||
effect_a,
|
||
M,
|
||
(box_w, box_h),
|
||
flags=cv2.INTER_LINEAR,
|
||
borderMode=cv2.BORDER_CONSTANT,
|
||
borderValue=0,
|
||
)
|
||
|
||
# 限制只在三角形内部混合(防止 bbox 多余像素污染)
|
||
tri_mask = np.zeros((box_h, box_w), dtype=np.uint8)
|
||
cv2.fillConvexPoly(
|
||
tri_mask,
|
||
dst_tri_local.astype(np.int32),
|
||
255,
|
||
lineType=cv2.LINE_AA,
|
||
)
|
||
alpha = (
|
||
warped_a.astype(np.float32) * (tri_mask.astype(np.float32) / 255.0)
|
||
) / 255.0
|
||
alpha = alpha[..., None] # (h,w,1)
|
||
|
||
roi = canvas_rgb[y_min_c:y_max_c, x_min_c:x_max_c].astype(np.float32)
|
||
blended = roi * (1.0 - alpha) + warped_rgb.astype(np.float32) * alpha
|
||
canvas_rgb[y_min_c:y_max_c, x_min_c:x_max_c] = np.clip(
|
||
blended, 0, 255
|
||
).astype(np.uint8)
|
||
|
||
def _encode_png(self, img_rgb: np.ndarray) -> bytes:
|
||
pil = Image.fromarray(img_rgb)
|
||
buf = io.BytesIO()
|
||
pil.save(buf, format="PNG")
|
||
return buf.getvalue()
|