andro1241 commited on
Commit
c8c1ed5
·
verified ·
1 Parent(s): 1ac698f

Upload 4 files

Browse files
processors/modules/expression_restorer/choices.py ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import List, Sequence, get_args
2
+
3
+ from facefusion.common_helper import create_int_range
4
+ from facefusion.processors.modules.expression_restorer.types import ExpressionRestorerArea, ExpressionRestorerModel
5
+
6
+ expression_restorer_models : List[ExpressionRestorerModel] = list(get_args(ExpressionRestorerModel))
7
+
8
+ expression_restorer_areas : List[ExpressionRestorerArea] = list(get_args(ExpressionRestorerArea))
9
+
10
+ expression_restorer_factor_range : Sequence[int] = create_int_range(0, 100, 1)
processors/modules/expression_restorer/core.py ADDED
@@ -0,0 +1,280 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from argparse import ArgumentParser
2
+ from functools import lru_cache
3
+ from types import ModuleType
4
+ from typing import List, Tuple
5
+
6
+ import cv2
7
+ import numpy
8
+
9
+ import facefusion.jobs.job_manager
10
+ import facefusion.jobs.job_store
11
+ from facefusion import config, content_analyser, face_classifier, face_detector, face_landmarker, face_masker, face_recognizer, inference_manager, logger, state_manager, translator, video_manager
12
+ from facefusion.common_helper import create_int_metavar, get_middle
13
+ from facefusion.download import conditional_download_hashes, conditional_download_sources, resolve_download_url
14
+ from facefusion.face_creator import scale_face
15
+ from facefusion.face_helper import paste_back, warp_face_by_face_landmark_5
16
+ from facefusion.face_masker import create_box_mask, create_occlusion_mask
17
+ from facefusion.face_selector import select_faces
18
+ from facefusion.filesystem import in_directory, is_image, is_video, resolve_relative_path, same_file_extension
19
+ from facefusion.processors.live_portrait import create_rotation, limit_expression
20
+ from facefusion.processors.modules.expression_restorer import choices as expression_restorer_choices
21
+ from facefusion.processors.modules.expression_restorer.types import ExpressionRestorerInputs
22
+ from facefusion.processors.types import LivePortraitExpression, LivePortraitFeatureVolume, LivePortraitMotionPoints, LivePortraitPitch, LivePortraitRoll, LivePortraitScale, LivePortraitTranslation, LivePortraitYaw, ProcessorOutputs
23
+ from facefusion.program_helper import find_argument_group
24
+ from facefusion.thread_helper import conditional_thread_semaphore, thread_semaphore
25
+ from facefusion.types import ApplyStateItem, Args, DownloadScope, Face, InferencePool, ModelOptions, ModelSet, ProcessMode, VisionFrame
26
+ from facefusion.vision import read_static_image, read_static_video_frame
27
+
28
+
29
+ @lru_cache()
30
+ def create_static_model_set(download_scope : DownloadScope) -> ModelSet:
31
+ return\
32
+ {
33
+ 'live_portrait':
34
+ {
35
+ '__metadata__':
36
+ {
37
+ 'vendor': 'KwaiVGI',
38
+ 'license': 'MIT',
39
+ 'year': 2024
40
+ },
41
+ 'hashes':
42
+ {
43
+ 'feature_extractor':
44
+ {
45
+ 'url': resolve_download_url('models-3.0.0', 'live_portrait_feature_extractor.hash'),
46
+ 'path': resolve_relative_path('../.assets/models/live_portrait_feature_extractor.hash')
47
+ },
48
+ 'motion_extractor':
49
+ {
50
+ 'url': resolve_download_url('models-3.0.0', 'live_portrait_motion_extractor.hash'),
51
+ 'path': resolve_relative_path('../.assets/models/live_portrait_motion_extractor.hash')
52
+ },
53
+ 'generator':
54
+ {
55
+ 'url': resolve_download_url('models-3.0.0', 'live_portrait_generator.hash'),
56
+ 'path': resolve_relative_path('../.assets/models/live_portrait_generator.hash')
57
+ }
58
+ },
59
+ 'sources':
60
+ {
61
+ 'feature_extractor':
62
+ {
63
+ 'url': resolve_download_url('models-3.0.0', 'live_portrait_feature_extractor.onnx'),
64
+ 'path': resolve_relative_path('../.assets/models/live_portrait_feature_extractor.onnx')
65
+ },
66
+ 'motion_extractor':
67
+ {
68
+ 'url': resolve_download_url('models-3.0.0', 'live_portrait_motion_extractor.onnx'),
69
+ 'path': resolve_relative_path('../.assets/models/live_portrait_motion_extractor.onnx')
70
+ },
71
+ 'generator':
72
+ {
73
+ 'url': resolve_download_url('models-3.0.0', 'live_portrait_generator.onnx'),
74
+ 'path': resolve_relative_path('../.assets/models/live_portrait_generator.onnx')
75
+ }
76
+ },
77
+ 'template': 'arcface_128',
78
+ 'size': (512, 512)
79
+ }
80
+ }
81
+
82
+
83
+ def get_inference_pool() -> InferencePool:
84
+ model_names = [ state_manager.get_item('expression_restorer_model') ]
85
+ model_source_set = get_model_options().get('sources')
86
+
87
+ return inference_manager.get_inference_pool(__name__, model_names, model_source_set)
88
+
89
+
90
+ def clear_inference_pool() -> None:
91
+ model_names = [ state_manager.get_item('expression_restorer_model') ]
92
+ inference_manager.clear_inference_pool(__name__, model_names)
93
+
94
+
95
+ def get_model_options() -> ModelOptions:
96
+ model_name = state_manager.get_item('expression_restorer_model')
97
+ return create_static_model_set('full').get(model_name)
98
+
99
+
100
+ def register_args(program : ArgumentParser) -> None:
101
+ group_processors = find_argument_group(program, 'processors')
102
+ if group_processors:
103
+ group_processors.add_argument('--expression-restorer-model', help = translator.get('help.model', __package__), default = config.get_str_value('processors', 'expression_restorer_model', 'live_portrait'), choices = expression_restorer_choices.expression_restorer_models)
104
+ group_processors.add_argument('--expression-restorer-factor', help = translator.get('help.factor', __package__), type = int, default = config.get_int_value('processors', 'expression_restorer_factor', '80'), choices = expression_restorer_choices.expression_restorer_factor_range, metavar = create_int_metavar(expression_restorer_choices.expression_restorer_factor_range))
105
+ group_processors.add_argument('--expression-restorer-areas', help = translator.get('help.areas', __package__).format(choices = ', '.join(expression_restorer_choices.expression_restorer_areas)), default = config.get_str_list('processors', 'expression_restorer_areas', ' '.join(expression_restorer_choices.expression_restorer_areas)), choices = expression_restorer_choices.expression_restorer_areas, nargs = '+', metavar = 'EXPRESSION_RESTORER_AREAS')
106
+ facefusion.jobs.job_store.register_step_keys([ 'expression_restorer_model', 'expression_restorer_factor', 'expression_restorer_areas' ])
107
+
108
+
109
+ def apply_args(args : Args, apply_state_item : ApplyStateItem) -> None:
110
+ apply_state_item('expression_restorer_model', args.get('expression_restorer_model'))
111
+ apply_state_item('expression_restorer_factor', args.get('expression_restorer_factor'))
112
+ apply_state_item('expression_restorer_areas', args.get('expression_restorer_areas'))
113
+
114
+
115
+ def get_common_modules() -> List[ModuleType]:
116
+ return [ content_analyser, face_classifier, face_detector, face_landmarker, face_masker, face_recognizer ]
117
+
118
+
119
+ def pre_check() -> bool:
120
+ model_hash_set = get_model_options().get('hashes')
121
+ model_source_set = get_model_options().get('sources')
122
+
123
+ for common_module in get_common_modules():
124
+ if not common_module.pre_check():
125
+ return False
126
+
127
+ return conditional_download_hashes(model_hash_set) and conditional_download_sources(model_source_set)
128
+
129
+
130
+ def pre_process(mode : ProcessMode) -> bool:
131
+ if mode == 'stream':
132
+ logger.error(translator.get('stream_not_supported') + translator.get('exclamation_mark'), __name__)
133
+ return False
134
+ if mode in [ 'output', 'preview' ] and not is_image(state_manager.get_item('target_path')) and not is_video(state_manager.get_item('target_path')):
135
+ logger.error(translator.get('choose_image_or_video_target') + translator.get('exclamation_mark'), __name__)
136
+ return False
137
+ if mode == 'output' and not in_directory(state_manager.get_item('output_path')):
138
+ logger.error(translator.get('specify_image_or_video_output') + translator.get('exclamation_mark'), __name__)
139
+ return False
140
+ if mode == 'output' and not same_file_extension(state_manager.get_item('target_path'), state_manager.get_item('output_path')):
141
+ logger.error(translator.get('match_target_and_output_extension') + translator.get('exclamation_mark'), __name__)
142
+ return False
143
+ return True
144
+
145
+
146
+ def post_process() -> None:
147
+ read_static_image.cache_clear()
148
+ read_static_video_frame.cache_clear()
149
+ video_manager.clear_video_pool()
150
+
151
+ if state_manager.get_item('video_memory_strategy') in [ 'strict', 'moderate' ]:
152
+ clear_inference_pool()
153
+
154
+ if state_manager.get_item('video_memory_strategy') == 'strict':
155
+ for common_module in get_common_modules():
156
+ common_module.clear_inference_pool()
157
+
158
+
159
+ def restore_expression(target_face : Face, target_vision_frame : VisionFrame, temp_vision_frame : VisionFrame) -> VisionFrame:
160
+ model_template = get_model_options().get('template')
161
+ model_size = get_model_options().get('size')
162
+ expression_restorer_factor = float(numpy.interp(float(state_manager.get_item('expression_restorer_factor')), [ 0, 100 ], [ 0, 1.2 ]))
163
+ target_crop_vision_frame, _ = warp_face_by_face_landmark_5(target_vision_frame, target_face.landmark_set.get('5/68'), model_template, model_size)
164
+ temp_crop_vision_frame, affine_matrix = warp_face_by_face_landmark_5(temp_vision_frame, target_face.landmark_set.get('5/68'), model_template, model_size)
165
+ box_mask = create_box_mask(temp_crop_vision_frame, state_manager.get_item('face_mask_blur'), (0, 0, 0, 0))
166
+ crop_masks =\
167
+ [
168
+ box_mask
169
+ ]
170
+
171
+ if 'occlusion' in state_manager.get_item('face_mask_types'):
172
+ occlusion_mask = create_occlusion_mask(temp_crop_vision_frame)
173
+ crop_masks.append(occlusion_mask)
174
+
175
+ target_crop_vision_frame = prepare_crop_frame(target_crop_vision_frame)
176
+ temp_crop_vision_frame = prepare_crop_frame(temp_crop_vision_frame)
177
+ temp_crop_vision_frame = apply_restore(target_crop_vision_frame, temp_crop_vision_frame, expression_restorer_factor)
178
+ temp_crop_vision_frame = normalize_crop_frame(temp_crop_vision_frame)
179
+ crop_mask = numpy.minimum.reduce(crop_masks).clip(0, 1)
180
+ paste_vision_frame = paste_back(temp_vision_frame, temp_crop_vision_frame, crop_mask, affine_matrix)
181
+ return paste_vision_frame
182
+
183
+
184
+ def apply_restore(target_crop_vision_frame : VisionFrame, temp_crop_vision_frame : VisionFrame, expression_restorer_factor : float) -> VisionFrame:
185
+ feature_volume = forward_extract_feature(temp_crop_vision_frame)
186
+ target_expression = forward_extract_motion(target_crop_vision_frame)[5]
187
+ pitch, yaw, roll, scale, translation, temp_expression, motion_points = forward_extract_motion(temp_crop_vision_frame)
188
+ rotation = create_rotation(pitch, yaw, roll)
189
+ target_expression = restrict_expression_areas(temp_expression, target_expression)
190
+ target_expression = target_expression * expression_restorer_factor + temp_expression * (1 - expression_restorer_factor)
191
+ target_expression = limit_expression(target_expression)
192
+ target_motion_points = scale * (motion_points @ rotation.T + target_expression) + translation
193
+ temp_motion_points = scale * (motion_points @ rotation.T + temp_expression) + translation
194
+ crop_vision_frame = forward_generate_frame(feature_volume, target_motion_points, temp_motion_points)
195
+ return crop_vision_frame
196
+
197
+
198
+ def restrict_expression_areas(temp_expression : LivePortraitExpression, target_expression : LivePortraitExpression) -> LivePortraitExpression:
199
+ expression_restorer_areas = state_manager.get_item('expression_restorer_areas')
200
+
201
+ if 'upper-face' not in expression_restorer_areas:
202
+ target_expression[:, [ 1, 2, 6, 10, 11, 12, 13, 15, 16 ]] = temp_expression[:, [ 1, 2, 6, 10, 11, 12, 13, 15, 16 ]]
203
+
204
+ if 'lower-face' not in expression_restorer_areas:
205
+ target_expression[:, [ 3, 7, 14, 17, 18, 19, 20 ]] = temp_expression[:, [ 3, 7, 14, 17, 18, 19, 20 ]]
206
+
207
+ target_expression[:, [ 0, 4, 5, 8, 9 ]] = temp_expression[:, [ 0, 4, 5, 8, 9 ]]
208
+ return target_expression
209
+
210
+
211
+ def forward_extract_feature(crop_vision_frame : VisionFrame) -> LivePortraitFeatureVolume:
212
+ feature_extractor = get_inference_pool().get('feature_extractor')
213
+
214
+ with conditional_thread_semaphore():
215
+ feature_volume = feature_extractor.run(None,
216
+ {
217
+ 'input': crop_vision_frame
218
+ })[0]
219
+
220
+ return feature_volume
221
+
222
+
223
+ def forward_extract_motion(crop_vision_frame : VisionFrame) -> Tuple[LivePortraitPitch, LivePortraitYaw, LivePortraitRoll, LivePortraitScale, LivePortraitTranslation, LivePortraitExpression, LivePortraitMotionPoints]:
224
+ motion_extractor = get_inference_pool().get('motion_extractor')
225
+
226
+ with conditional_thread_semaphore():
227
+ pitch, yaw, roll, scale, translation, expression, motion_points = motion_extractor.run(None,
228
+ {
229
+ 'input': crop_vision_frame
230
+ })
231
+
232
+ return pitch, yaw, roll, scale, translation, expression, motion_points
233
+
234
+
235
+ def forward_generate_frame(feature_volume : LivePortraitFeatureVolume, target_motion_points : LivePortraitMotionPoints, temp_motion_points : LivePortraitMotionPoints) -> VisionFrame:
236
+ generator = get_inference_pool().get('generator')
237
+
238
+ with thread_semaphore():
239
+ crop_vision_frame = generator.run(None,
240
+ {
241
+ 'feature_volume': feature_volume,
242
+ 'source': target_motion_points,
243
+ 'target': temp_motion_points
244
+ })[0][0]
245
+
246
+ return crop_vision_frame
247
+
248
+
249
+ def prepare_crop_frame(crop_vision_frame : VisionFrame) -> VisionFrame:
250
+ model_size = get_model_options().get('size')
251
+ prepare_size = (model_size[0] // 2, model_size[1] // 2)
252
+ crop_vision_frame = cv2.resize(crop_vision_frame, prepare_size, interpolation = cv2.INTER_AREA)
253
+ crop_vision_frame = crop_vision_frame[:, :, ::-1] / 255.0
254
+ crop_vision_frame = numpy.expand_dims(crop_vision_frame.transpose(2, 0, 1), axis = 0).astype(numpy.float32)
255
+ return crop_vision_frame
256
+
257
+
258
+ def normalize_crop_frame(crop_vision_frame : VisionFrame) -> VisionFrame:
259
+ crop_vision_frame = crop_vision_frame.transpose(1, 2, 0).clip(0, 1)
260
+ crop_vision_frame = crop_vision_frame * 255.0
261
+ crop_vision_frame = crop_vision_frame.astype(numpy.uint8)[:, :, ::-1]
262
+ return crop_vision_frame
263
+
264
+
265
+ def process_frame(inputs : ExpressionRestorerInputs) -> ProcessorOutputs:
266
+ reference_vision_frame = inputs.get('reference_vision_frame')
267
+ source_vision_frames = inputs.get('source_vision_frames')
268
+ target_vision_frames = inputs.get('target_vision_frames')
269
+ temp_vision_frame = inputs.get('temp_vision_frame')
270
+ temp_vision_mask = inputs.get('temp_vision_mask')
271
+
272
+ target_vision_frame = get_middle(target_vision_frames)
273
+ target_faces = select_faces(reference_vision_frame, source_vision_frames, target_vision_frames)
274
+
275
+ if target_faces:
276
+ for target_face in target_faces:
277
+ target_face = scale_face(target_face, target_vision_frame, temp_vision_frame)
278
+ temp_vision_frame = restore_expression(target_face, target_vision_frame, temp_vision_frame)
279
+
280
+ return temp_vision_frame, temp_vision_mask
processors/modules/expression_restorer/locales.py ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from facefusion.types import Locales
2
+
3
+ LOCALES : Locales =\
4
+ {
5
+ 'en':
6
+ {
7
+ 'help':
8
+ {
9
+ 'model': 'choose the model responsible for restoring the expression',
10
+ 'factor': 'restore factor of expression from the target face',
11
+ 'areas': 'choose the items used for the expression areas (choices: {choices})'
12
+ },
13
+ 'uis':
14
+ {
15
+ 'model_dropdown': 'EXPRESSION RESTORER MODEL',
16
+ 'factor_slider': 'EXPRESSION RESTORER FACTOR',
17
+ 'areas_checkbox_group': 'EXPRESSION RESTORER AREAS'
18
+ }
19
+ }
20
+ }
processors/modules/expression_restorer/types.py ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import List, Literal, TypedDict
2
+
3
+ from facefusion.types import Mask, VisionFrame
4
+
5
+ ExpressionRestorerInputs = TypedDict('ExpressionRestorerInputs',
6
+ {
7
+ 'reference_vision_frame' : VisionFrame,
8
+ 'source_vision_frames' : List[VisionFrame],
9
+ 'target_vision_frames' : List[VisionFrame],
10
+ 'temp_vision_frame' : VisionFrame,
11
+ 'temp_vision_mask' : Mask
12
+ })
13
+
14
+ ExpressionRestorerModel = Literal['live_portrait']
15
+
16
+ ExpressionRestorerArea = Literal['upper-face', 'lower-face']