@@ -90,6 +90,14 @@ class HairStyle_Model(object):
self . keypoints_processor = hair_init . keypoints_processor
self . human_keypoint = hair_init . human_keypoint
# 用户图预处理产物内存缓存(一级缓存,避免重复 GPU 推理 + 磁盘 IO)。
# key = 图片字节哈希 + ratio; value = 7 个产物 + landmarks_1k 的 dict。
# 接口2 女性多发型场景:同一张用户图连续请求,第二次起命中缓存省 ~1.6s。
import collections
self . _user_prepare_cache = { } # {cache_key: {产物dict}}
self . _user_prepare_cache_keys = collections . deque ( ) # LRU 顺序
self . _USER_CACHE_MAX = 8 # 最多缓存 8 张图(防止内存膨胀)
# worker_id = int(os.environ.get('APP_WORKER_ID', 1))
# rand_max = 9527
@@ -2268,48 +2276,99 @@ class HairStyle_Model(object):
ref_landmark_f1k2_768 )
cv2 . imwrite ( another_pose_hair_image_dir , another_pose_hair_image * 255 )
landmark1k_dir = osp . join ( userinfo_dir , ' kpt_1k.txt ' )
if not osp . exists ( landmark1k_dir ) :
landmarks_origin_img_1k , bounding_box , euler_info = self . get_landmark . forward_diy ( user_rgb_8uc3_orisize )
# ===== 内存缓存(一级):key = 图片哈希 + ratio =====
# 命中则跳过 landmark 检测 + get_prepare_user_768_data 整条 GPU 管线(省 ~1.6s)
import hashlib as _hashlib
_img_hash = _hashlib . md5 ( user_rgb_8uc3_orisize . tobytes ( ) ) . hexdigest ( ) [ : 16 ]
_cache_key = f " { _img_hash } _r { ratio } "
_cached = self . _user_prepare_cache . get ( _cache_key )
if _cached is not None :
# 命中内存缓存:直接取所有产物,跳过 landmark 检测和预处理管线
landmarks_origin_img_1k = _cached [ ' landmarks_1k ' ]
user_bald_res_8uc3_orisize = _cached [ ' bald_res ' ]
user_baldseg_8uc3_orisize = _cached [ ' baldseg_ori ' ]
user_baldseg_8uc3_768 = _cached [ ' baldseg_768 ' ]
user_bald_8uc3_768 = _cached [ ' bald_768 ' ]
user_landmark_f1k2_768 = _cached [ ' lmk_768 ' ]
user_hairstyle_M = _cached [ ' M ' ]
user_matting_8uc3_bald_orisize = _cached [ ' matting_ori ' ]
# 补写磁盘文件:下游代码(功能7等)仍从 userinfo_dir 读这些文件,
# 而 task_id 每次不同导致 userinfo_dir 不同,必须补写保证下游可用
os . makedirs ( userinfo_dir , exist_ok = True )
np . savetxt ( landmark1k_dir , landmarks_origin_img_1k )
cv2 . imwrite ( osp . join ( userinfo_dir , ' bald_res_ori.png ' ) , user_bald_res_8uc3_orisize )
cv2 . imwrite ( osp . join ( userinfo_dir , ' bald_seg_ori.png ' ) , user_baldseg_8uc3_orisize )
cv2 . imwrite ( osp . join ( userinfo_dir , ' user_baldseg_768.png ' ) , user_baldseg_8uc3_768 )
cv2 . imwrite ( osp . join ( userinfo_dir , ' bald_seg_768.png ' ) , user_bald_8uc3_768 )
np . savetxt ( osp . join ( userinfo_dir , ' landmark_f1k2_768.txt ' ) , user_landmark_f1k2_768 )
np . savetxt ( osp . join ( userinfo_dir , ' hairstyle_M.txt ' ) , user_hairstyle_M )
if user_matting_8uc3_bald_orisize is not None :
cv2 . imwrite ( osp . join ( userinfo_dir , ' user_orig_mask.png ' ) , user_matting_8uc3_bald_orisize )
self . logger_process . info ( f " 内存缓存命中 key= { _cache_key } ,跳过用户图预处理(补写磁盘文件) " )
else :
# 未命中:走原有逻辑(landmark 检测 + 预处理管线 + 磁盘缓存)
if not osp . exists ( landmark1k_dir ) :
landmarks_origin_img_1k , bounding_box , euler_info = self . get_landmark . forward_diy ( user_rgb_8uc3_orisize )
if landmarks_origin_img_1k is None :
return None , 10001 , None , None , None
np . savetxt ( landmark1k_dir , landmarks_origin_img_1k )
else :
landmarks_origin_img_1k = np . loadtxt ( landmark1k_dir )
# landmarks_origin_img_1k, _, _ = self.get_landmark.forward(user_rgb_8uc3_orisize)
if landmarks_origin_img_1k is None :
return None , 10001 , None , None , None
np . savetxt ( landmark1k_dir , landmarks_origin_img_1k )
else :
landmarks_origin_img_1k = np . loadtxt ( landmark1k_dir )
# landmarks_origin_img_1k, _, _ = self.get_landmark.forward(user_rgb_8uc3_orisize)
if landmarks_origin_img_1k is None :
return None , 10001 , None , None , None
user_bald_res_8uc3_orisize_dir = osp . join ( userinfo_dir , ' bald_res_ori.png ' )
user_baldseg_8uc3_orisize_dir = osp . join ( userinfo_dir , ' bald_seg_ori.png ' )
user_baldseg_8uc3_768_dir = osp . join ( userinfo_dir , ' user_baldseg_768.png ' )
user_bald_8uc3_768_dir = osp . join ( userinfo_dir , ' bald_seg_768.png ' )
user_landmark_f1k2_768_dir = osp . join ( userinfo_dir , ' landmark_f1k2_768.txt ' )
user_hairstyle_M_dir = osp . join ( userinfo_dir , ' hairstyle_M.txt ' )
user_bald_res_8uc3_orisize_dir = osp . join ( userinfo_dir , ' bald_res_ori.png ' )
user_baldseg_8uc3_orisize_dir = osp . join ( userinfo_dir , ' bald_seg_ori.png ' )
user_baldseg_8uc3_768_dir = osp . join ( userinfo_dir , ' user_baldseg_768.png ' )
user_bald_8uc3_768_dir = osp . join ( userinfo_dir , ' bald_seg_768.png ' )
user_landmark_f1k2_768_dir = osp . join ( userinfo_dir , ' landmark_f1k2_768.txt ' )
user_hairstyle_M_dir = osp . join ( userinfo_dir , ' hairstyle_M.txt ' )
condition_exist2 = True
pre_list = [ user_bald_res_8uc3_orisize_dir , user_baldseg_8uc3_orisize_dir , user_baldseg_8uc3_768_dir , user_bald_8uc3_768_dir ,
user_landmark_f1k2_768_dir , user_hairstyle_M_dir ]
for tmp_dir in pre_list :
if not osp . exists ( tmp_dir ) :
condition_exist2 = False
condition_exist2 = True
pre_list = [ user_bald_res_8uc3_orisize_dir , user_baldseg_8uc3_orisize_dir , user_baldseg_8uc3_768_dir , user_bald_8uc3_768_dir ,
user_landmark_f1k2_768_dir , user_hairstyle_M_dir ]
for tmp_dir in pre_list :
if not osp . exists ( tmp_dir ) :
condition_exist2 = False
user_matting_8uc3_bald_orisize = None
if not condition_exist2 :
user_bald_res_8uc3_orisize , user_baldseg_8uc3_orisize , user_baldseg_8uc3_768 , user_bald_8uc3_768 , \
user_landmark_f1k2_768 , user_hairstyle_M , user_matting_8uc3_bald_orisize = self . process_data . get_prepare_user_768_data ( user_rgb_8uc3_orisize , landmarks_origin_img_1k , ratio = ratio )
cv2 . imwrite ( user_bald_res_8uc3_orisize_dir , user_bald_res_8uc3_orisize )
cv2 . imwrite ( user_baldseg_8uc3_orisize_dir , user_baldseg_8uc3_orisize )
cv2 . imwrite ( user_baldseg_8uc3_768_dir , user_baldseg_8uc3_768 )
cv2 . imwrite ( user_bald_8uc3_768_dir , user_bald_8uc3_768 )
np . savetxt ( user_landmark_f1k2_768_dir , user_landmark_f1k2_768 )
np . savetxt ( user_hairstyle_M_dir , user_hairstyle_M )
else :
user_bald_res_8uc3_orisize = cv2 . imread ( user_bald_res_8uc3_orisize_dir )
user_baldseg_8uc3_orisize = cv2 . imread ( user_baldseg_8uc3_orisize_dir )
user_baldseg_8uc3_768 = cv2 . imread ( user_baldseg_8uc3_768_dir )
user_bald_8uc3_768 = cv2 . imread ( user_bald_8uc3_768_dir )
user_landmark_f1k2_768 = np . loadtxt ( user_landmark_f1k2_768_dir )
user_hairstyle_M = np . loadtxt ( user_hairstyle_M_dir )
user_matting_8uc3_bald_orisize = None
if not condition_exist2 :
user_bald_res_8uc3_orisize , user_baldseg_8uc3_orisize , user_baldseg_8uc3_768 , user_bald_8uc3_768 , \
user_landmark_f1k2_768 , user_hairstyle_M , user_matting_8uc3_bald_orisize = self . process_data . get_prepare_user_768_data ( user_rgb_8uc3_orisize , landmarks_origin_img_1k , ratio = ratio )
cv2 . imwrite ( user_bald_res_8uc3_orisize_dir , user_bald_res_8uc3_orisize )
cv2 . imwrite ( user_baldseg_8uc3_orisize_dir , user_baldseg_8uc3_orisize )
cv2 . imwrite ( user_baldseg_8uc3_768_dir , user_baldseg_8uc3_768 )
cv2 . imwrite ( user_bald_8uc3_768_dir , user_bald_8uc3_768 )
np . savetxt ( user_landmark_f1k2_768_dir , user_landmark_f1k2_768 )
np . savetxt ( user_hairstyle_M_dir , user_hairstyle_M )
else :
user_bald_res_8uc3_orisize = cv2 . imread ( user_bald_res_8uc3_orisize_dir )
user_baldseg_8uc3_orisize = cv2 . imread ( user_baldseg_8uc3_orisize_dir )
user_baldseg_8uc3_768 = cv2 . imread ( user_baldseg_8uc3_768_dir )
user_bald_8uc3_768 = cv2 . imread ( user_bald_8uc3_768_dir )
user_landmark_f1k2_768 = np . loadtxt ( user_landmark_f1k2_768_dir )
user_hairstyle_M = np . loadtxt ( user_hairstyle_M_dir )
# 写入内存缓存(含 landmarks_1k 和 matting,供后续同图请求命中)
self . _user_prepare_cache [ _cache_key ] = {
' landmarks_1k ' : landmarks_origin_img_1k ,
' bald_res ' : user_bald_res_8uc3_orisize ,
' baldseg_ori ' : user_baldseg_8uc3_orisize ,
' baldseg_768 ' : user_baldseg_8uc3_768 ,
' bald_768 ' : user_bald_8uc3_768 ,
' lmk_768 ' : user_landmark_f1k2_768 ,
' M ' : user_hairstyle_M ,
' matting_ori ' : user_matting_8uc3_bald_orisize ,
}
self . _user_prepare_cache_keys . append ( _cache_key )
# LRU 淘汰:超过上限删除最老的
while len ( self . _user_prepare_cache_keys ) > self . _USER_CACHE_MAX :
_old = self . _user_prepare_cache_keys . popleft ( )
self . _user_prepare_cache . pop ( _old , None )
self . logger_process . info ( f " 内存缓存写入 key= { _cache_key } ,当前缓存 { len ( self . _user_prepare_cache ) } 张 " )
# show_concat = np.concatenate((user_rgb_8uc3_orisize, user_bald_res_8uc3_orisize, user_baldseg_8uc3_orisize), axis=1)
# ratio = 1536. / max(show_concat.shape[:2])
@@ -2318,18 +2377,26 @@ class HairStyle_Model(object):
# cv2.waitKey()
user_orig_mask_path = os . path . join ( userinfo_dir , " user_orig_mask.png " )
if not os . path . exists ( user_orig_mask_path ) :
if user_matting_8uc3_bald_orisize is not None and not os . path . exists ( user_orig_mask_path ) :
cv2 . imwrite ( user_orig_mask_path , user_matting_8uc3_bald_orisize )
# 换发型
# 换发型(主 GAN)
import time as _time
_t_gan0 = _time . perf_counter ( )
hair_gene_8uc3_768 = self . generator_hair . Generator_Hair_inference_use_pref ( another_pose_hair_image ,
user_baldseg_8uc3_768 ,
user_bald_8uc3_768 ,
user_landmark_f1k2_768 , gender )
_t_gan = _time . perf_counter ( ) - _t_gan0
# 融合(第3次matte + 融合GAN)
_t_fuse0 = _time . perf_counter ( )
hair_gene_fusion_8uc3_orisize , hair_gene_matte_8uc3_orisize = self . process_data . get_fusion_res_hairpaste (
user_bald_res_8uc3_orisize , hair_gene_8uc3_768 , user_landmark_f1k2_768 , user_hairstyle_M )
_t_fuse = _time . perf_counter ( ) - _t_fuse0
self . logger_process . info (
f " 功能6分步计时: 主GAN= { _t_gan : .3f } s 融合= { _t_fuse : .3f } s (缓存= { ' 命中 ' if _cached is not None else ' 未命中 ' } ) " )
gen_hair_mask_path_2 = os . path . join ( userinfo_dir , " hair_mask_2.png " )
cv2 . imwrite ( gen_hair_mask_path_2 , hair_gene_matte_8uc3_orisize )