Files
change_hair/hair_grow_cli.py
xsl 6f34e8876c feat: 适配 RTX 3090 (24GB) 环境优化
硬件迁移:从 RTX 5090 (32GB) 迁移到 RTX 3090 (24GB)

主要改动:
- 所有启动脚本和配置文件中的 /home/xsl/ 路径替换为 /home/ubuntu/
- 适配新的 24GB VRAM 环境
2026-07-18 19:00:15 +08:00

90 lines
3.2 KiB
Python
Raw Permalink 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/ubuntu/change_hair/project/hair_service_sd
/home/ubuntu/miniconda3/envs/my_hair/bin/python /home/ubuntu/change_hair/hair_grow_cli.py \
--img test.jpg --mask mask.png --strength 0.5 -o result.jpg
# 批量跑三档强度对比
/home/ubuntu/miniconda3/envs/my_hair/bin/python /home/ubuntu/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/ubuntu/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()