初始化:换发型/换发色/训练发型服务
包含: - 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密钥已脱敏为环境变量,原文件备份在本地
This commit is contained in:
@@ -0,0 +1,89 @@
|
||||
#!/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()
|
||||
Reference in New Issue
Block a user