Files
change_hair/hair_grow_cli.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

90 lines
3.1 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_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()