Files
change_hair/gen_template_material.py
T
xsl 443cfa298f 初始化:换发型/换发色/训练发型服务
包含:
- hair_service_sd: 主服务(换发型/换发色/生发,端口8801)
- photo_service: LoRA调度+训练(端口32678)
- hair_grow_service: 调试测试页(端口8888,含4个测试页)
- 批量训练脚本(batch_train_hairstyles.py)
- 发际线mask自动识别(hairline_mask.py,4种方案)
- 手绘mask换发型(hair_swap_manual.py)
- 文档:README.md + LARGE_FILES.md + docs/

大文件(模型权重200G、训练数据123G)已排除,见 LARGE_FILES.md
OSS/COS密钥已脱敏为环境变量,原文件备份在本地
2026-07-07 13:53:52 +08:00

100 lines
3.3 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.
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""为新发型生成 hair_template_materialfirst##_matting.png + .pkl
换发型功能4 需要:
hair_template_material/<hair_id>/first##<name>.png 原图
hair_template_material/<hair_id>/first##<name>_matting.png 头发抠图mask
hair_template_material/<hair_id>/first##<name>.pkl 1k关键点
用法:
cd /home/xsl/change_hair/project/hair_service_sd
python gen_template_material.py --hair-id new_test_001 --img /path/to/template.jpg
"""
import os
import sys
import ssl
import uuid
import pickle
import argparse
ssl._create_default_https_context = ssl._create_unverified_context
HAIR_SERVICE_DIR = "/home/xsl/change_hair/project/hair_service_sd"
os.chdir(HAIR_SERVICE_DIR)
sys.path.insert(0, HAIR_SERVICE_DIR)
os.environ.setdefault("HF_HUB_OFFLINE", "1")
import cv2
import numpy as np
import torch
_orig_torch_load = torch.load
def _patched_torch_load(*args, **kwargs):
if 'map_location' in kwargs and callable(kwargs['map_location']) and not isinstance(kwargs['map_location'], str):
kwargs['map_location'] = 'cpu'
return _orig_torch_load(*args, **kwargs)
torch.load = _patched_torch_load
from models.detector import RetinaFaceDetector
from models.MomocvFaceAlignment1K import MomocvFaceAlignment1K
from utils.landmark_processor import get_max_rect
from core.process_modules import Generator_Matte
from common.logger import config
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--hair-id", required=True)
parser.add_argument("--img", required=True, help="模板图(发型图)")
args = parser.parse_args()
template_dir = config.get('default', 'hair_template_material_dir')
out_dir = os.path.join(template_dir, args.hair_id)
os.makedirs(out_dir, exist_ok=True)
img = cv2.imread(args.img)
if img is None:
print(f"✗ 读取失败: {args.img}"); sys.exit(1)
print(f"模板图: {img.shape}")
# 1. 人脸检测 + 1k关键点
print("加载模型...")
detector = RetinaFaceDetector(gpu_id=0)
aligner = MomocvFaceAlignment1K(gpu_id=0)
matte_gen = Generator_Matte(gpu=True, device_id=0)
print("模型加载完成")
dets, landms = detector.forward(img)
if len(dets) == 0:
print("✗ 未检测到人脸"); sys.exit(1)
det = dets[get_max_rect(dets)]
landmarks_1k = aligner.stable_forward(img.copy(), [det])
pt1k = landmarks_1k[0]
# 2. 头发抠图
with torch.no_grad():
_, alpha, _ = matte_gen.matte_inference(img, pt1k)
if alpha.shape != img.shape[:2]:
alpha = cv2.resize(alpha, (img.shape[1], img.shape[0]))
# 3. 保存三个文件(first##前缀)
file_uuid = f"first##{uuid.uuid4()}"
png_path = os.path.join(out_dir, f"{file_uuid}.png")
matting_path = os.path.join(out_dir, f"{file_uuid}_matting.png")
pkl_path = os.path.join(out_dir, f"{file_uuid}.pkl")
cv2.imwrite(png_path, img)
cv2.imwrite(matting_path, alpha)
with open(pkl_path, 'wb') as f:
# 存 numpy array(不能用 tolist(),代码用 landmarks[index, :] 索引需要 ndarray
pickle.dump({'human_pt1k': np.asarray(pt1k, dtype=np.float32)}, f)
print(f"✓ 生成完成:")
print(f" {png_path}")
print(f" {matting_path}")
print(f" {pkl_path}")
if __name__ == "__main__":
main()