包含: - 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密钥已脱敏为环境变量,原文件备份在本地
90 lines
3.1 KiB
Python
90 lines
3.1 KiB
Python
#!/usr/bin/env python3
|
||
# -*- coding: utf-8 -*-
|
||
"""区域生发 - 独立验证脚本
|
||
|
||
不依赖 hair_service 主服务,直接调用 hair_grow.hair_grow()。
|
||
需要在 hair_service_sd 目录下运行(依赖其模块导入)。
|
||
|
||
用法:
|
||
cd /home/xsl/change_hair/project/hair_service_sd
|
||
/home/xsl/miniconda3/envs/my_hair/bin/python /home/xsl/change_hair/hair_grow_cli.py \
|
||
--img test.jpg --mask mask.png --strength 0.5 -o result.jpg
|
||
|
||
# 批量跑三档强度对比
|
||
/home/xsl/miniconda3/envs/my_hair/bin/python /home/xsl/change_hair/hair_grow_cli.py \
|
||
--img test.jpg --mask mask.png --compare
|
||
"""
|
||
import os
|
||
import sys
|
||
import cv2
|
||
import argparse
|
||
|
||
# 把 hair_service_sd 加入路径,使其模块可被导入
|
||
HAIR_SERVICE_DIR = "/home/xsl/change_hair/project/hair_service_sd"
|
||
sys.path.insert(0, HAIR_SERVICE_DIR)
|
||
|
||
# 设置离线模式(本机无法访问 huggingface.co)
|
||
os.environ.setdefault("HF_HUB_OFFLINE", "1")
|
||
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
|
||
|
||
|
||
def run_once(img_path, mask_path, strength, out_path):
|
||
"""单次生发"""
|
||
from hair_grow import hair_grow
|
||
|
||
img = cv2.imread(img_path)
|
||
mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)
|
||
if img is None:
|
||
print(f"❌ 无法读取图片: {img_path}")
|
||
return False
|
||
if mask is None:
|
||
print(f"❌ 无法读取mask: {mask_path}")
|
||
return False
|
||
|
||
print(f"输入: img={img.shape}, mask={mask.shape}, strength={strength}")
|
||
result = hair_grow(img, mask, strength=strength)
|
||
|
||
cv2.imwrite(out_path, result)
|
||
print(f"✅ 生发结果已保存: {out_path}")
|
||
return True
|
||
|
||
|
||
def run_compare(img_path, mask_path, out_dir):
|
||
"""三档强度对比"""
|
||
os.makedirs(out_dir, exist_ok=True)
|
||
for strength in [0.2, 0.5, 0.8]:
|
||
out_path = os.path.join(out_dir, f"result_s{strength}.jpg")
|
||
print(f"\n{'='*50}")
|
||
print(f"生发强度 strength={strength}")
|
||
print(f"{'='*50}")
|
||
if not run_once(img_path, mask_path, strength, out_path):
|
||
return False
|
||
print(f"\n✅ 三档对比完成,结果在: {out_dir}")
|
||
print(" - result_s0.2.jpg (低强度,轻微生发)")
|
||
print(" - result_s0.5.jpg (中强度,明显生发)")
|
||
print(" - result_s0.8.jpg (高强度,浓密生发)")
|
||
return True
|
||
|
||
|
||
def main():
|
||
parser = argparse.ArgumentParser(description="区域生发验证脚本")
|
||
parser.add_argument("--img", required=True, help="人头像图片路径")
|
||
parser.add_argument("--mask", required=True, help="遮罩图片路径(白=生发区)")
|
||
parser.add_argument("--strength", type=float, default=0.5,
|
||
help="生发强度 0.1~1.0 (默认0.5)")
|
||
parser.add_argument("-o", "--output", default="result_hairgrow.jpg",
|
||
help="输出图片路径")
|
||
parser.add_argument("--compare", action="store_true",
|
||
help="跑三档强度(0.2/0.5/0.8)对比")
|
||
args = parser.parse_args()
|
||
|
||
if args.compare:
|
||
out_dir = os.path.dirname(args.output) or "."
|
||
run_compare(args.img, args.mask, out_dir)
|
||
else:
|
||
run_once(args.img, args.mask, args.strength, args.output)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|