fix: 适配 ubuntu 路径并增强换发色/训练流程稳定性
将配置与训练脚本从 /home/xsl 切到本机 /home/ubuntu;换发色在 webui 增强失败或缺色板时降级返回,训练结束后自动重启 hair 服务再回调。 Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -70,7 +70,11 @@ def process_infer(user_img_path, target_color, color_dir, dst_path, ratio=1.0):
|
||||
# 通过webui对照片进行增强
|
||||
s4 = time.time()
|
||||
# cv2.imwrite('/home/student/Downloads/12121.jpg',crop_img)
|
||||
enhanced_img = enhance_hair.webui_img2img(crop_img, crop_mask, prompt=prompt)
|
||||
try:
|
||||
enhanced_img = enhance_hair.webui_img2img(crop_img, crop_mask, prompt=prompt)
|
||||
except Exception as e:
|
||||
print(f"[process_infer] webui enhance skipped: {e}")
|
||||
enhanced_img = crop_img
|
||||
print("----------------------------webui_img2img", time.time() - s4)
|
||||
|
||||
# 写回原图
|
||||
|
||||
@@ -1,21 +1,21 @@
|
||||
[default]
|
||||
modelDir = weights
|
||||
hairstyleDir = /home/xsl/change_hair/project/data/ref_hairstyle
|
||||
haircolorDir= /home/xsl/change_hair/project/data/ref_haircolor
|
||||
userDir=/home/xsl/change_hair/project/data/userImage
|
||||
tmp_dir=/home/xsl/change_hair/project/data/tmp
|
||||
res_dir=/home/xsl/change_hair/project/data/res_dir
|
||||
userInfo_dir=/home/xsl/change_hair/project/data/user_info
|
||||
hairstyleDir = /home/ubuntu/change_hair/project/data/ref_hairstyle
|
||||
haircolorDir= /home/ubuntu/change_hair/project/data/ref_haircolor
|
||||
userDir=/home/ubuntu/change_hair/project/data/userImage
|
||||
tmp_dir=/home/ubuntu/change_hair/project/data/tmp
|
||||
res_dir=/home/ubuntu/change_hair/project/data/res_dir
|
||||
userInfo_dir=/home/ubuntu/change_hair/project/data/user_info
|
||||
baseColor_ID=HDR10_443322
|
||||
Port = 11023
|
||||
refer_dir = /home/xsl/change_hair/project/data/ref_online
|
||||
ref_user_dir = /home/xsl/change_hair/project/data/ref_user_imgs
|
||||
train_dir = /home/xsl/change_hair/data/train_material
|
||||
refImgDir=/home/xsl/change_hair/project/data/refImage
|
||||
hair_template_material_dir=/home/xsl/change_hair/project/data/hair_template_material
|
||||
ref_color=/home/xsl/change_hair/project/data/ref_color
|
||||
ref_color_img=/home/xsl/change_hair/project/data/ref_color_imgs
|
||||
upload_train_dir=/home/xsl/change_hair/project/data/upload_train_imgs
|
||||
refer_dir = /home/ubuntu/change_hair/project/data/ref_online
|
||||
ref_user_dir = /home/ubuntu/change_hair/project/data/ref_user_imgs
|
||||
train_dir = /home/ubuntu/change_hair/data/train_material
|
||||
refImgDir=/home/ubuntu/change_hair/project/data/refImage
|
||||
hair_template_material_dir=/home/ubuntu/change_hair/project/data/hair_template_material
|
||||
ref_color=/home/ubuntu/change_hair/project/data/ref_color
|
||||
ref_color_img=/home/ubuntu/change_hair/project/data/ref_color_imgs
|
||||
upload_train_dir=/home/ubuntu/change_hair/project/data/upload_train_imgs
|
||||
;version=local
|
||||
version=online
|
||||
|
||||
@@ -23,7 +23,7 @@ version=online
|
||||
strength=1
|
||||
|
||||
[logger]
|
||||
logpath = /home/xsl/change_hair/project/logs
|
||||
logpath = /home/ubuntu/change_hair/project/logs
|
||||
level=INFO
|
||||
|
||||
[timelogger]
|
||||
@@ -32,4 +32,3 @@ name=watch-time
|
||||
[errorlogger]
|
||||
level=ERROR
|
||||
name=w-error
|
||||
|
||||
|
||||
@@ -72,6 +72,8 @@ class HairStyle_Model_Infer(object):
|
||||
# haircolor_dir = '/home/data/hair/data/ref_color/3628746832766'
|
||||
face_base, hair_matting, status = self.infer_haircolor_new(user_rgb_8uc3_orisize, haircolor_dir,target_hair_color,
|
||||
return_matting=True)
|
||||
if status != 0:
|
||||
return user_rgb_8uc3_orisize, None, status
|
||||
user_rgb_8uc3_orisize = face_base
|
||||
# 构建一个色板
|
||||
r, g, b = target_hair_color
|
||||
@@ -1016,7 +1018,7 @@ class HairStyle_Model_Infer(object):
|
||||
def infer_haircolor_new(self, user_rgb_8uc3_orisize, haircolor_dir, target_hair_color, return_matting=False):
|
||||
landmarks_origin_img_1k= self.get_landmark.forward(user_rgb_8uc3_orisize)
|
||||
if landmarks_origin_img_1k is None:
|
||||
return None, 10001
|
||||
return None, None, 10001
|
||||
need_process = (target_hair_color[0] * 0.299 + target_hair_color[1] * 0.587 + target_hair_color[2] * 0.114) > 150
|
||||
_, user_matting_8uc1_bald_orisize = self.process_data_infer.generator_matte.matte_inference(user_rgb_8uc3_orisize,
|
||||
landmarks_origin_img_1k)
|
||||
@@ -1036,10 +1038,13 @@ class HairStyle_Model_Infer(object):
|
||||
need_process=True
|
||||
if need_process:
|
||||
haircolor_dir_tmp = os.path.join(config.get('default', "haircolorDir"), config.get('default', "baseColor_ID"))
|
||||
face_base, status1 = self.infer_haircolor_tj(user_rgb_8uc3_orisize, haircolor_dir_tmp)
|
||||
# cv2.imwrite('/home/student/Desktop/tmp_color/need/face_base.png', face_base)
|
||||
if status1 == 0:
|
||||
return face_base, user_matting_8uc1_bald_orisize, 0
|
||||
if os.path.exists(haircolor_dir_tmp):
|
||||
face_base, status1 = self.infer_haircolor_tj(user_rgb_8uc3_orisize, haircolor_dir_tmp)
|
||||
# cv2.imwrite('/home/student/Desktop/tmp_color/need/face_base.png', face_base)
|
||||
if status1 == 0:
|
||||
return face_base, user_matting_8uc1_bald_orisize, 0
|
||||
else:
|
||||
need_process = False
|
||||
else:
|
||||
need_process = False
|
||||
if not need_process:
|
||||
|
||||
@@ -255,7 +255,8 @@ def change_hair_colorv3():
|
||||
jsonify({'msg': '算法解析错误', 'result': '', 'umd': '', 'state': -1}), 400)
|
||||
|
||||
except Exception as e:
|
||||
print(e)
|
||||
print(f"[hairColor ERROR] {e}")
|
||||
traceback.print_exc()
|
||||
return make_response(
|
||||
jsonify({'msg': '算法解析错误', 'result': '', 'umd': '', 'state': -1}), 400)
|
||||
|
||||
|
||||
@@ -43,6 +43,9 @@ def webui_img2img(img, mask, prompt='', denoising_strength=0.35):
|
||||
}
|
||||
response = requests.post(url=url, json=request_dict)
|
||||
ret_json = response.json()
|
||||
if 'images' not in ret_json:
|
||||
print(f"[webui_img2img ERROR] {ret_json}")
|
||||
raise Exception(f"webui error: {ret_json.get('error', 'unknown')}")
|
||||
result = ret_json['images'][0]
|
||||
img = cv2.imdecode(np.frombuffer(base64.b64decode(result.split(",", 1)[0]), np.uint8), cv2.IMREAD_COLOR)
|
||||
return img
|
||||
|
||||
@@ -37,8 +37,8 @@ if version == "local":
|
||||
callback_url = 'http://service.aicloud.fit:7395/api/hair/trainCallBack'
|
||||
else:
|
||||
current_url = 'http://0.0.0.0:7393/'
|
||||
kohya_ss_home_dir = '/home/xsl/change_hair/project/kohya_ss_home'
|
||||
webui_lora_dir = '/home/xsl/change_hair/project/onediff/stable-diffusion-webui/models/Lora'
|
||||
kohya_ss_home_dir = '/home/ubuntu/change_hair/project/kohya_ss_home'
|
||||
webui_lora_dir = '/home/ubuntu/change_hair/project/onediff/stable-diffusion-webui/models/Lora'
|
||||
inference_use_onediff = False
|
||||
callback_url = 'http://0.0.0.0:8801/api/hair/trainCallBack'
|
||||
base_webui_port = '57860'
|
||||
@@ -307,13 +307,13 @@ def train_thread(sq, gpu_id):
|
||||
# 3. GPU 固定 device=0(单卡)
|
||||
# 4. 去掉 tokenizer_cache_dir(改用 HF 本地缓存 + 离线模式)
|
||||
# 5. 设置 HF_HUB_OFFLINE 避免联网检查
|
||||
kohya_python = '/home/xsl/miniconda3/envs/kohya/bin/python'
|
||||
kohya_python = '/home/ubuntu/miniconda3/envs/my_hair/bin/python'
|
||||
kohya_workdir = os.path.join(kohya_ss_home_dir, 'kohya_ss')
|
||||
base_model = '/home/xsl/change_hair/project/onediff/stable-diffusion-webui/models/Stable-diffusion/v1-5-pruned-emaonly.safetensors'
|
||||
base_model = '/home/ubuntu/change_hair/project/onediff/stable-diffusion-webui/models/Stable-diffusion/v1-5-pruned-emaonly.safetensors'
|
||||
cmd_train = (
|
||||
f'cd {kohya_workdir} && '
|
||||
f'HF_HUB_OFFLINE=1 TRANSFORMERS_OFFLINE=1 CUDA_VISIBLE_DEVICES={device_id} '
|
||||
f'/home/xsl/miniconda3/envs/kohya/bin/accelerate launch --num_cpu_threads_per_process=2 "./train_network.py" --enable_bucket '
|
||||
f'/home/ubuntu/miniconda3/envs/my_hair/bin/accelerate launch --num_cpu_threads_per_process=2 "./train_network.py" --enable_bucket '
|
||||
f'--min_bucket_reso=256 --max_bucket_reso=2048 --pretrained_model_name_or_path="{base_model}" '
|
||||
f'--train_data_dir={images_dir} --resolution="2000,2000" '
|
||||
f'--output_dir={model_dir} '
|
||||
|
||||
@@ -158,6 +158,10 @@ def step3_wait_and_callback(hair_id, template_img):
|
||||
else:
|
||||
print(" ✗ 训练超时"); return False
|
||||
time.sleep(5) # 等训练进程写完
|
||||
# 训练时停掉了 hair 服务,回调前需要重启
|
||||
print(" 重启 change_hair-hair 服务...")
|
||||
os.system("sudo systemctl restart change_hair-hair")
|
||||
time.sleep(20)
|
||||
# 触发回调生成ref材质
|
||||
print(" 触发回调生成ref材质...")
|
||||
r = req.post(HAIR_CALLBACK, json={"task_id":f"train_{hair_id}","hair_id":hair_id,
|
||||
|
||||
Reference in New Issue
Block a user