innoai commited on
Commit
915d52b
·
verified ·
1 Parent(s): 985635d

Update src/gradio_pipeline.py

Browse files
Files changed (1) hide show
  1. src/gradio_pipeline.py +174 -62
src/gradio_pipeline.py CHANGED
@@ -2,8 +2,29 @@
2
 
3
  """
4
  Pipeline for gradio
 
 
 
 
 
5
  """
 
6
  import gradio as gr
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7
  from .config.argument_config import ArgumentConfig
8
  from .live_portrait_pipeline import LivePortraitPipeline
9
  from .utils.io import load_img_online
@@ -12,21 +33,25 @@ from .utils.crop import prepare_paste_back, paste_back
12
  from .utils.camera import get_rotation_matrix
13
  from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio
14
 
 
15
  def update_args(args, user_args):
16
- """update the args according to user inputs
 
17
  """
18
  for k, v in user_args.items():
19
  if hasattr(args, k):
20
  setattr(args, k, v)
21
  return args
22
 
 
23
  class GradioPipeline(LivePortraitPipeline):
24
 
25
  def __init__(self, inference_cfg, crop_cfg, args: ArgumentConfig):
26
  super().__init__(inference_cfg, crop_cfg)
27
- # self.live_portrait_wrapper = self.live_portrait_wrapper
28
  self.args = args
29
- # for single image retargeting
 
30
  self.start_prepare = False
31
  self.f_s_user = None
32
  self.x_c_s_info_user = None
@@ -35,8 +60,9 @@ class GradioPipeline(LivePortraitPipeline):
35
  self.mask_ori = None
36
  self.img_rgb = None
37
  self.crop_M_c2o = None
 
38
 
39
-
40
  def execute_video(
41
  self,
42
  input_image_path,
@@ -44,99 +70,185 @@ class GradioPipeline(LivePortraitPipeline):
44
  flag_relative_input,
45
  flag_do_crop_input,
46
  flag_remap_input,
47
- ):
48
- """ for video driven potrait animation
 
 
 
 
49
  """
50
  if input_image_path is not None and input_video_path is not None:
51
  args_user = {
52
- 'source_image': input_image_path,
53
- 'driving_info': input_video_path,
54
- 'flag_relative': flag_relative_input,
55
- 'flag_do_crop': flag_do_crop_input,
56
- 'flag_pasteback': flag_remap_input,
57
  }
58
- # update config from user input
 
59
  self.args = update_args(self.args, args_user)
60
  self.live_portrait_wrapper.update_config(self.args.__dict__)
61
  self.cropper.update_config(self.args.__dict__)
62
- # video driven animation
 
63
  video_path, video_path_concat = self.execute(self.args)
64
- # gr.Info("Run successfully!", duration=2)
65
- return video_path, video_path_concat,
66
- else:
67
- raise gr.Error("The input source portrait or driving video hasn't been prepared yet 💥!", duration=5)
68
 
 
 
 
 
 
 
 
 
69
  def execute_image(self, input_eye_ratio: float, input_lip_ratio: float):
70
- """ for single image retargeting
71
  """
72
- if input_eye_ratio is None or input_eye_ratio is None:
 
 
 
 
 
 
73
  raise gr.Error("Invalid ratio input 💥!", duration=5)
74
- elif self.f_s_user is None:
 
75
  if self.start_prepare:
76
  raise gr.Error(
77
  "The source portrait is under processing 💥! Please wait for a second.",
78
- duration=5
79
- )
80
- else:
81
- raise gr.Error(
82
- "The source portrait hasn't been prepared yet 💥! Please scroll to the top of the page to upload.",
83
- duration=5
84
  )
85
- else:
86
- x_s_user = self.x_s_user.to("cuda")
87
- f_s_user = self.f_s_user.to("cuda")
88
- # ∆_eyes,i = R_eyes(x_s; c_s,eyes, c_d,eyes,i)
89
- combined_eye_ratio_tensor = self.live_portrait_wrapper.calc_combined_eye_ratio([[input_eye_ratio]], self.source_lmk_user)
90
- eyes_delta = self.live_portrait_wrapper.retarget_eye(x_s_user, combined_eye_ratio_tensor)
91
- # ∆_lip,i = R_lip(x_s; c_s,lip, c_d,lip,i)
92
- combined_lip_ratio_tensor = self.live_portrait_wrapper.calc_combined_lip_ratio([[input_lip_ratio]], self.source_lmk_user)
93
- lip_delta = self.live_portrait_wrapper.retarget_lip(x_s_user, combined_lip_ratio_tensor)
94
- num_kp = x_s_user.shape[1]
95
- # default: use x_s
96
- x_d_new = x_s_user + eyes_delta.reshape(-1, num_kp, 3) + lip_delta.reshape(-1, num_kp, 3)
97
- # D(W(f_s; x_s, x′_d))
98
- out = self.live_portrait_wrapper.warp_decode(f_s_user, x_s_user, x_d_new)
99
- out = self.live_portrait_wrapper.parse_output(out['out'])[0]
100
- out_to_ori_blend = paste_back(out, self.crop_M_c2o, self.img_rgb, self.mask_ori)
101
- # gr.Info("Run successfully!", duration=2)
102
- return out, out_to_ori_blend
103
-
104
-
105
- def prepare_retargeting(self, input_image_path, flag_do_crop = True):
106
- """ for single image retargeting
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
107
  """
108
  if input_image_path is not None:
109
- # gr.Info("Upload successfully!", duration=2)
110
  self.start_prepare = True
 
111
  inference_cfg = self.live_portrait_wrapper.cfg
112
- ######## process source portrait ########
113
- img_rgb = load_img_online(input_image_path, mode='rgb', max_dim=1280, n=16)
 
 
 
 
 
 
114
  log(f"Load source image from {input_image_path}.")
 
 
115
  crop_info = self.cropper.crop_single_image(img_rgb)
 
116
  if flag_do_crop:
117
- I_s = self.live_portrait_wrapper.prepare_source(crop_info['img_crop_256x256'])
 
 
118
  else:
119
  I_s = self.live_portrait_wrapper.prepare_source(img_rgb)
 
 
120
  x_s_info = self.live_portrait_wrapper.get_kp_info(I_s)
121
- R_s = get_rotation_matrix(x_s_info['pitch'], x_s_info['yaw'], x_s_info['roll'])
122
- ############################################
123
 
124
- # record global info for next time use
 
 
 
 
 
 
 
 
125
  self.f_s_user = self.live_portrait_wrapper.extract_feature_3d(I_s)
126
  self.x_s_user = self.live_portrait_wrapper.transform_keypoint(x_s_info)
127
  self.x_s_info_user = x_s_info
128
- self.source_lmk_user = crop_info['lmk_crop']
129
  self.img_rgb = img_rgb
130
- self.crop_M_c2o = crop_info['M_c2o']
131
- self.mask_ori = prepare_paste_back(inference_cfg.mask_crop, crop_info['M_c2o'], dsize=(img_rgb.shape[1], img_rgb.shape[0]))
132
- # update slider
 
 
 
 
 
133
  eye_close_ratio = calc_eye_close_ratio(self.source_lmk_user[None])
134
  eye_close_ratio = float(eye_close_ratio.squeeze(0).mean())
 
135
  lip_close_ratio = calc_lip_close_ratio(self.source_lmk_user[None])
136
  lip_close_ratio = float(lip_close_ratio.squeeze(0).mean())
137
- # for vis
 
138
  self.I_s_vis = self.live_portrait_wrapper.parse_output(I_s)[0]
 
 
 
139
  return eye_close_ratio, lip_close_ratio, self.I_s_vis
140
- else:
141
- # when press the clear button, go here
 
142
  return 0.8, 0.8, self.I_s_vis
 
 
 
2
 
3
  """
4
  Pipeline for gradio
5
+
6
+ 已适配 Hugging Face ZeroGPU:
7
+ 1. 给会使用 CUDA 的 Gradio 回调函数添加 @spaces.GPU 装饰器。
8
+ 2. 保留 Gradio 4.x 兼容写法。
9
+ 3. 如果本地环境没有 spaces 包,也不会影响本地普通运行。
10
  """
11
+
12
  import gradio as gr
13
+
14
+ # Hugging Face ZeroGPU 必须使用 spaces.GPU 装饰需要 GPU 的函数。
15
+ # 为了兼容本地运行,如果没有安装 spaces,则使用一个空装饰器。
16
+ try:
17
+ import spaces
18
+ except Exception:
19
+ class _DummySpaces:
20
+ @staticmethod
21
+ def GPU(duration=120):
22
+ def decorator(func):
23
+ return func
24
+ return decorator
25
+
26
+ spaces = _DummySpaces()
27
+
28
  from .config.argument_config import ArgumentConfig
29
  from .live_portrait_pipeline import LivePortraitPipeline
30
  from .utils.io import load_img_online
 
33
  from .utils.camera import get_rotation_matrix
34
  from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio
35
 
36
+
37
  def update_args(args, user_args):
38
+ """
39
+ 根据用户输入更新参数。
40
  """
41
  for k, v in user_args.items():
42
  if hasattr(args, k):
43
  setattr(args, k, v)
44
  return args
45
 
46
+
47
  class GradioPipeline(LivePortraitPipeline):
48
 
49
  def __init__(self, inference_cfg, crop_cfg, args: ArgumentConfig):
50
  super().__init__(inference_cfg, crop_cfg)
51
+
52
  self.args = args
53
+
54
+ # 单图重定向状态缓存
55
  self.start_prepare = False
56
  self.f_s_user = None
57
  self.x_c_s_info_user = None
 
60
  self.mask_ori = None
61
  self.img_rgb = None
62
  self.crop_M_c2o = None
63
+ self.I_s_vis = None
64
 
65
+ @spaces.GPU(duration=300)
66
  def execute_video(
67
  self,
68
  input_image_path,
 
70
  flag_relative_input,
71
  flag_do_crop_input,
72
  flag_remap_input,
73
+ ):
74
+ """
75
+ 视频驱动肖像动画。
76
+
77
+ 注意:
78
+ 这个函数内部会调用模型推理和 CUDA,因此必须放在 @spaces.GPU 里。
79
  """
80
  if input_image_path is not None and input_video_path is not None:
81
  args_user = {
82
+ "source_image": input_image_path,
83
+ "driving_info": input_video_path,
84
+ "flag_relative": flag_relative_input,
85
+ "flag_do_crop": flag_do_crop_input,
86
+ "flag_pasteback": flag_remap_input,
87
  }
88
+
89
+ # 根据用户输入更新配置
90
  self.args = update_args(self.args, args_user)
91
  self.live_portrait_wrapper.update_config(self.args.__dict__)
92
  self.cropper.update_config(self.args.__dict__)
93
+
94
+ # 执行视频驱动动画
95
  video_path, video_path_concat = self.execute(self.args)
 
 
 
 
96
 
97
+ return video_path, video_path_concat
98
+
99
+ raise gr.Error(
100
+ "The input source portrait or driving video hasn't been prepared yet 💥!",
101
+ duration=5,
102
+ )
103
+
104
+ @spaces.GPU(duration=180)
105
  def execute_image(self, input_eye_ratio: float, input_lip_ratio: float):
 
106
  """
107
+ 单图表情重定向。
108
+
109
+ 注意:
110
+ 这里会执行 .to("cuda")、retarget_eye、retarget_lip、warp_decode,
111
+ 因此必须放在 @spaces.GPU 里。
112
+ """
113
+ if input_eye_ratio is None or input_lip_ratio is None:
114
  raise gr.Error("Invalid ratio input 💥!", duration=5)
115
+
116
+ if self.f_s_user is None:
117
  if self.start_prepare:
118
  raise gr.Error(
119
  "The source portrait is under processing 💥! Please wait for a second.",
120
+ duration=5,
 
 
 
 
 
121
  )
122
+
123
+ raise gr.Error(
124
+ "The source portrait hasn't been prepared yet 💥! Please scroll to the top of the page to upload.",
125
+ duration=5,
126
+ )
127
+
128
+ x_s_user = self.x_s_user.to("cuda")
129
+ f_s_user = self.f_s_user.to("cuda")
130
+
131
+ # 计算眼睛重定向
132
+ combined_eye_ratio_tensor = self.live_portrait_wrapper.calc_combined_eye_ratio(
133
+ [[input_eye_ratio]],
134
+ self.source_lmk_user,
135
+ )
136
+ eyes_delta = self.live_portrait_wrapper.retarget_eye(
137
+ x_s_user,
138
+ combined_eye_ratio_tensor,
139
+ )
140
+
141
+ # 计算嘴唇重定向
142
+ combined_lip_ratio_tensor = self.live_portrait_wrapper.calc_combined_lip_ratio(
143
+ [[input_lip_ratio]],
144
+ self.source_lmk_user,
145
+ )
146
+ lip_delta = self.live_portrait_wrapper.retarget_lip(
147
+ x_s_user,
148
+ combined_lip_ratio_tensor,
149
+ )
150
+
151
+ num_kp = x_s_user.shape[1]
152
+
153
+ # 默认基于 x_s 做变形
154
+ x_d_new = (
155
+ x_s_user
156
+ + eyes_delta.reshape(-1, num_kp, 3)
157
+ + lip_delta.reshape(-1, num_kp, 3)
158
+ )
159
+
160
+ # 解码输出
161
+ out = self.live_portrait_wrapper.warp_decode(
162
+ f_s_user,
163
+ x_s_user,
164
+ x_d_new,
165
+ )
166
+ out = self.live_portrait_wrapper.parse_output(out["out"])[0]
167
+
168
+ # 贴回原图
169
+ out_to_ori_blend = paste_back(
170
+ out,
171
+ self.crop_M_c2o,
172
+ self.img_rgb,
173
+ self.mask_ori,
174
+ )
175
+
176
+ return out, out_to_ori_blend
177
+
178
+ @spaces.GPU(duration=180)
179
+ def prepare_retargeting(self, input_image_path, flag_do_crop=True):
180
+ """
181
+ 单图表情重定向的预处理。
182
+
183
+ 注意:
184
+ 日志中的报错发生在这个函数调用链里:
185
+ prepare_retargeting -> prepare_source -> x.cuda(...)
186
+ 所以这个函数必须使用 @spaces.GPU。
187
  """
188
  if input_image_path is not None:
 
189
  self.start_prepare = True
190
+
191
  inference_cfg = self.live_portrait_wrapper.cfg
192
+
193
+ # 读取源图
194
+ img_rgb = load_img_online(
195
+ input_image_path,
196
+ mode="rgb",
197
+ max_dim=1280,
198
+ n=16,
199
+ )
200
  log(f"Load source image from {input_image_path}.")
201
+
202
+ # 裁剪人脸
203
  crop_info = self.cropper.crop_single_image(img_rgb)
204
+
205
  if flag_do_crop:
206
+ I_s = self.live_portrait_wrapper.prepare_source(
207
+ crop_info["img_crop_256x256"]
208
+ )
209
  else:
210
  I_s = self.live_portrait_wrapper.prepare_source(img_rgb)
211
+
212
+ # 提取关键点信息
213
  x_s_info = self.live_portrait_wrapper.get_kp_info(I_s)
 
 
214
 
215
+ # 保留原逻辑:计算旋转矩阵
216
+ # 当前变量暂未在后续使用,但保留,避免影响原项目行为。
217
+ _ = get_rotation_matrix(
218
+ x_s_info["pitch"],
219
+ x_s_info["yaw"],
220
+ x_s_info["roll"],
221
+ )
222
+
223
+ # 缓存后续单���重定向需要的数据
224
  self.f_s_user = self.live_portrait_wrapper.extract_feature_3d(I_s)
225
  self.x_s_user = self.live_portrait_wrapper.transform_keypoint(x_s_info)
226
  self.x_s_info_user = x_s_info
227
+ self.source_lmk_user = crop_info["lmk_crop"]
228
  self.img_rgb = img_rgb
229
+ self.crop_M_c2o = crop_info["M_c2o"]
230
+ self.mask_ori = prepare_paste_back(
231
+ inference_cfg.mask_crop,
232
+ crop_info["M_c2o"],
233
+ dsize=(img_rgb.shape[1], img_rgb.shape[0]),
234
+ )
235
+
236
+ # 更新滑块默认值
237
  eye_close_ratio = calc_eye_close_ratio(self.source_lmk_user[None])
238
  eye_close_ratio = float(eye_close_ratio.squeeze(0).mean())
239
+
240
  lip_close_ratio = calc_lip_close_ratio(self.source_lmk_user[None])
241
  lip_close_ratio = float(lip_close_ratio.squeeze(0).mean())
242
+
243
+ # 预览图
244
  self.I_s_vis = self.live_portrait_wrapper.parse_output(I_s)[0]
245
+
246
+ self.start_prepare = False
247
+
248
  return eye_close_ratio, lip_close_ratio, self.I_s_vis
249
+
250
+ # 点击清空按钮时走这里
251
+ if self.I_s_vis is not None:
252
  return 0.8, 0.8, self.I_s_vis
253
+
254
+ return 0.8, 0.8, None