部署修复: - torch.load 增加 weights_only=False patch,兼容 PyTorch 2.6+ 加载旧权重 - OSS 改为懒加载,本地用 output_format=base64 无需配凭证即可启动 - 补全被 gitignore 误排除的必需代码:core/models/layers/data、models/layers/data、keypoints/lib - webui 训练命令 --xformers 改 --sdpa(修复 xformers 无 CUDA 支持报错) 功能调整: - hair_grow_service 端口改 8899、preview 路由修复(send_file) - list_hairstyles 增加发型白名单,测试页只展示当前5个发型 新增脚本: - train_lora_parallel.py:直接调 kohya 并行训练 LoRA(绕过 photo_service 串行限制) - train_hairstyles_parallel.py / train_batch_stepC.py:批量训练辅助脚本 - scripts/sync_data_to_server.sh:大文件断点续传到云服务器 文档: - docs/换发型集成文档.md:换发型完整流程、服务架构、资源依赖、训练方法、集成步骤
205 lines
9.0 KiB
Python
205 lines
9.0 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""生发模块(走换发型工作流)
|
||
|
||
和换发型 /api/swapHair/v1 完全一致的工作流:infer_hairstyle_diy_jy 粗推理
|
||
→ warpAffine → photo_service + LoRA → webui inpainting → 贴回原图。
|
||
唯一区别:mask 是 origin_matting ∪ new_matting ∪ 手绘mask 的并集,
|
||
不减刘海/眼眶。
|
||
|
||
核心函数:
|
||
hair_grow_swap(origin_img, hand_mask, hair_id, hairstyle_process,
|
||
landmark_processor, task_id, is_hr) -> BGR ndarray
|
||
"""
|
||
import os
|
||
import json
|
||
import glob
|
||
import time
|
||
import shutil
|
||
|
||
import cv2
|
||
import numpy as np
|
||
|
||
# 复用换发型的工具链
|
||
from gen_super_image import webui_img2img
|
||
from common.logger import config
|
||
from utils.landmark_processor import high_quality_warpAffine
|
||
|
||
|
||
def hair_grow_swap(origin_img, hand_mask, hair_id, hairstyle_process,
|
||
landmark_processor, task_id, is_hr=True):
|
||
"""在用户涂抹的缺发区生发(走换发型工作流 + 该发型 LoRA)。
|
||
|
||
参数:
|
||
origin_img: 用户原图 BGR ndarray
|
||
hand_mask: 用户手绘mask,灰度 ndarray,白色(255)=生发区,与 origin_img 同尺寸
|
||
hair_id: 选择的发型ID(用于加载 LoRA + 读材质/prompt)
|
||
hairstyle_process: hair_init 初始化的 HairstyleModel 实例
|
||
landmark_processor: landmark_processor 模块(含 high_quality_warpAffine)
|
||
task_id: 任务ID(唯一字符串)
|
||
is_hr: 是否高清模式(True=1152x1536,False=576x768)
|
||
|
||
返回:
|
||
BGR ndarray 生发结果图(与原图同尺寸)
|
||
"""
|
||
start_all = time.time()
|
||
|
||
# ===== 路径配置(与 swapHair 一致)=====
|
||
hairstyle_dir = config.get('default', 'hairstyleDir') # ref_hairstyle/
|
||
user_dir = config.get('default', 'userDir')
|
||
train_dir = config.get('default', 'train_dir') # train_material/
|
||
userInfo_dir = config.get('default', 'userInfo_dir')
|
||
res_dir = config.get('default', 'res_dir')
|
||
|
||
hair_material_dir = os.path.join(train_dir, hair_id) # train_material/<id>
|
||
material_save_path = os.path.join(hairstyle_dir, hair_id) # ref_hairstyle/<id>
|
||
|
||
# 校验:ref 材质必须存在(infer_hairstyle_diy_jy 依赖)
|
||
ref_img_path = os.path.join(material_save_path, "ref_rgb_8uc3_768.png")
|
||
if not os.path.exists(ref_img_path):
|
||
raise FileNotFoundError(f"发型材质不存在: {ref_img_path}")
|
||
|
||
# 保存用户原图到工作目录(infer_hairstyle_diy_jy 需要读文件路径)
|
||
user_img_name = f"hairgrow_{task_id}.jpg"
|
||
new_user_img_path = os.path.join(user_dir, user_img_name)
|
||
if os.path.exists(new_user_img_path):
|
||
os.remove(new_user_img_path)
|
||
cv2.imwrite(new_user_img_path, origin_img)
|
||
origin_img = cv2.imread(new_user_img_path)
|
||
|
||
# ===== 功能6: infer_hairstyle_diy_jy 粗推理 =====
|
||
# 产出 img_res(带新发型的粗推理图)+ 临时文件(user_orig_mask.png, hair_mask_2.png, kpt_1k.txt)
|
||
start = time.time()
|
||
ref_img = cv2.imread(ref_img_path)
|
||
work_dir = os.path.join(userInfo_dir, task_id)
|
||
os.makedirs(work_dir, exist_ok=True)
|
||
import torch
|
||
with torch.no_grad():
|
||
img_res, status, _, landmarks_1k, isEyeOccluded = hairstyle_process.infer_hairstyle_diy_jy(
|
||
origin_img, ref_img, os.path.join(work_dir, task_id), f"{task_id}.png")
|
||
if status != 0:
|
||
raise RuntimeError("发型粗推理失败")
|
||
print(f"[hair_grow] 功能6 粗推理完成: {time.time()-start:.1f}s")
|
||
|
||
# ===== 功能8: warpAffine + mask 合并(核心改造点)=====
|
||
start = time.time()
|
||
user_material_dir = os.path.join(work_dir, task_id)
|
||
hair_matting_path = os.path.join(user_material_dir, "hair_mask_2.png")
|
||
user_orig_mask_path = os.path.join(user_material_dir, "user_orig_mask.png")
|
||
|
||
origin_matting = cv2.imread(user_orig_mask_path, cv2.IMREAD_GRAYSCALE)
|
||
new_matting = cv2.imread(hair_matting_path, cv2.IMREAD_GRAYSCALE)
|
||
if origin_matting is None or new_matting is None:
|
||
raise RuntimeError(f"粗推理产物缺失: origin={origin_matting is None}, new={new_matting is None}")
|
||
|
||
# 【关键差异】mask = origin_matting ∪ new_matting ∪ hand_mask
|
||
# - 不减刘海圆、不减眼眶(与换发型不同)
|
||
# - 手绘 mask 加入并集,覆盖用户指定的缺发区
|
||
# 手绘 mask 与 origin_matting 同尺寸(都是原图尺寸)
|
||
if hand_mask.shape != origin_matting.shape:
|
||
hand_mask = cv2.resize(hand_mask, (origin_matting.shape[1], origin_matting.shape[0]),
|
||
interpolation=cv2.INTER_NEAREST)
|
||
_, hand_mask_bin = cv2.threshold(hand_mask, 127, 255, cv2.THRESH_BINARY)
|
||
matting_merge = np.max(
|
||
np.stack([origin_matting, new_matting, hand_mask_bin], axis=2), axis=2
|
||
).astype(np.uint8)
|
||
|
||
# warpAffine 到固定尺寸(与换发型第514-527行完全一致)
|
||
dst_size = (1152, 1536) if is_hr else (576, 768)
|
||
box_info = hairstyle_process.get_body_info(img_res)
|
||
box_w = box_info[2] - box_info[0]
|
||
box_h = box_info[3] - box_info[1]
|
||
scale = min(dst_size[1] / max(box_h, 1), dst_size[0] / max(box_w, 1))
|
||
rotate_center = [(box_info[2] + box_info[0]) * 0.5, (box_info[3] + box_info[1]) * 0.5]
|
||
M = cv2.getRotationMatrix2D(rotate_center, 0, scale)
|
||
M[:, 2] += np.float32([dst_size[0] * 0.5, dst_size[1] * 0.5]) - np.float32(rotate_center)
|
||
|
||
crop_result = high_quality_warpAffine(img_res, M, dst_size)
|
||
crop_matting = cv2.warpAffine(matting_merge, M, dst_size)
|
||
mask = (crop_matting > 10).astype(np.float32)
|
||
# dilate(与换发型一致:非HR核3x9,HR核6x18)
|
||
dilate_kernel = (6, 18) if is_hr else (3, 9)
|
||
mask_dilate = cv2.dilate(mask, np.ones(dilate_kernel, np.uint8))
|
||
mask_dilate = np.clip(mask_dilate * 255, 0, 255).astype(np.uint8)
|
||
final_img = crop_result
|
||
|
||
# 读取 config.json 的 gender + 训练素材的 prompt tag(与换发型第565-578行一致)
|
||
config_json_path = os.path.join(material_save_path, "config.json")
|
||
with open(config_json_path, "r") as f:
|
||
in_gender = json.load(f)["gender"]
|
||
|
||
images_dir = os.path.join(hair_material_dir, "images")
|
||
txt_dir = os.path.join(images_dir, os.listdir(images_dir)[0])
|
||
txt_path = glob.glob(txt_dir + '/*.txt')[0]
|
||
with open(txt_path, 'r') as f:
|
||
p_tag = f.readline()
|
||
if "titor hairstyle, faceless, no human, gray background, simple background" in p_tag:
|
||
p_tag = p_tag[p_tag.find("simple background, ") + len("simple background, "):]
|
||
else:
|
||
p_tag = ""
|
||
print(f"[hair_grow] 功能8 mask合并+warpAffine完成: {time.time()-start:.1f}s, "
|
||
f"in_gender={in_gender}")
|
||
|
||
# ===== 功能9: webui 推理(走 photo_service + LoRA,复用 webui_img2img)=====
|
||
start = time.time()
|
||
sd_result = webui_img2img(
|
||
img=final_img, mask_img=mask_dilate,
|
||
in_gender=in_gender, task_id=task_id,
|
||
hair_id=hair_id, lora_material_path=hair_material_dir,
|
||
tag=p_tag, is_hr=is_hr,
|
||
denoising_strength=0.6, inference_port="57860")
|
||
print(f"[hair_grow] 功能9 webui推理完成: {time.time()-start:.1f}s")
|
||
|
||
# ===== 后处理: warpAffine 贴回原图(与换发型第601-605行一致)=====
|
||
M_inv = cv2.invertAffineTransform(M)
|
||
result = origin_img.copy()
|
||
cv2.warpAffine(sd_result, M_inv, (origin_img.shape[1], origin_img.shape[0]),
|
||
dst=result, borderMode=cv2.BORDER_TRANSPARENT, flags=cv2.INTER_LANCZOS4)
|
||
|
||
# 清理临时文件
|
||
try:
|
||
shutil.rmtree(user_material_dir)
|
||
os.remove(new_user_img_path)
|
||
except Exception:
|
||
pass
|
||
|
||
print(f"[hair_grow] 总耗时: {time.time()-start_all:.1f}s")
|
||
return result
|
||
|
||
|
||
def list_hairstyles(hairstyle_dir, train_dir, upload_dir, limit=None):
|
||
"""列出所有完整可用的发型。
|
||
|
||
参数:
|
||
hairstyle_dir: ref_hairstyle/ 目录
|
||
train_dir: train_material/ 目录
|
||
upload_dir: upload_train_imgs/ 目录
|
||
limit: 限制返回数量(测试用)
|
||
返回:
|
||
[{hair_id, gender, has_lora}, ...]
|
||
"""
|
||
# 仅展示这些发型(其余数据保留在磁盘,但不在列表中显示)
|
||
_VISIBLE_HAIRSTYLES = {
|
||
"chang_tuoyuan", "chang_bolang", "chang_zhixian",
|
||
"chang_huaban", "chang_xinxing",
|
||
}
|
||
styles = []
|
||
if not os.path.isdir(hairstyle_dir):
|
||
return styles
|
||
for hair_id in os.listdir(hairstyle_dir):
|
||
if hair_id not in _VISIBLE_HAIRSTYLES:
|
||
continue
|
||
cfg_path = os.path.join(hairstyle_dir, hair_id, "config.json")
|
||
ref_path = os.path.join(hairstyle_dir, hair_id, "ref_rgb_8uc3_768.png")
|
||
lora_path = os.path.join(train_dir, hair_id, "model", "hairstyle_hd_lora.safetensors")
|
||
if not (os.path.exists(cfg_path) and os.path.exists(ref_path) and os.path.exists(lora_path)):
|
||
continue
|
||
try:
|
||
with open(cfg_path) as f:
|
||
gender = json.load(f).get("gender", "unknown")
|
||
except Exception:
|
||
gender = "unknown"
|
||
styles.append({"hair_id": hair_id, "gender": gender})
|
||
if limit and len(styles) >= limit:
|
||
break
|
||
return styles
|